use std::collections::HashMap;
use std::io::Cursor;
use std::path::{Path, PathBuf};
#[cfg(feature = "fetch-images")]
use std::sync::Mutex;
use docling_core::PictureImage;
#[cfg(feature = "fetch-images")]
const MAX_IMAGE_BYTES: u64 = 32 * 1024 * 1024;
pub(crate) trait ImageResolver {
fn resolve(&self, src: &str) -> Option<PictureImage>;
fn prefetch(&self, _srcs: &[String]) {}
}
pub(crate) struct NoFetch;
impl ImageResolver for NoFetch {
fn resolve(&self, _src: &str) -> Option<PictureImage> {
None
}
}
pub(crate) struct FsImageResolver {
base_dir: Option<PathBuf>,
base_url: Option<String>,
#[cfg(feature = "fetch-images")]
cache: Mutex<HashMap<String, Option<PictureImage>>>,
}
impl FsImageResolver {
pub(crate) fn new(base_dir: Option<PathBuf>, base_url: Option<String>) -> Self {
Self {
base_dir,
base_url,
#[cfg(feature = "fetch-images")]
cache: Mutex::new(HashMap::new()),
}
}
#[cfg(feature = "fetch-images")]
fn absolute_http_url(&self, src: &str) -> Option<String> {
if src.starts_with("http://") || src.starts_with("https://") {
return Some(src.to_string());
}
let base = self.base_url.as_deref()?;
let joined = url::Url::parse(base).ok()?.join(src).ok()?;
matches!(joined.scheme(), "http" | "https").then(|| joined.to_string())
}
#[cfg(feature = "fetch-images")]
fn fetch_cached(&self, url: &str) -> Option<PictureImage> {
if let Some(hit) = self.cache.lock().unwrap().get(url) {
return hit.clone();
}
let img = fetch_remote(url);
self.cache
.lock()
.unwrap()
.insert(url.to_string(), img.clone());
img
}
}
#[cfg(feature = "fetch-images")]
fn image_fetch_concurrency() -> usize {
std::env::var("DOCLING_RS_IMAGE_FETCH_CONCURRENCY")
.ok()
.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|&n| n > 0)
.unwrap_or(10)
.clamp(1, 64)
}
impl ImageResolver for FsImageResolver {
fn resolve(&self, src: &str) -> Option<PictureImage> {
let src = src.trim();
if src.is_empty() {
return None;
}
if src.starts_with("data:") {
return from_data_uri(src);
}
#[cfg(feature = "fetch-images")]
if let Some(url) = self.absolute_http_url(src) {
return self.fetch_cached(&url);
}
#[cfg(not(feature = "fetch-images"))]
if src.starts_with("http://") || src.starts_with("https://") {
return None;
}
let rel = src.strip_prefix("file://").unwrap_or(src);
let path = Path::new(rel);
let full = if path.is_absolute() {
path.to_path_buf()
} else {
self.base_dir.as_ref()?.join(rel)
};
let data = std::fs::read(&full).ok()?;
super::ooxml::picture_image(full.to_str().unwrap_or(rel), data)
}
#[cfg(feature = "fetch-images")]
fn prefetch(&self, srcs: &[String]) {
let urls: Vec<String> = {
let cache = self.cache.lock().unwrap();
let mut seen = std::collections::HashSet::new();
srcs.iter()
.filter_map(|s| self.absolute_http_url(s.trim()))
.filter(|u| !cache.contains_key(u) && seen.insert(u.clone()))
.collect()
};
if urls.len() < 2 {
return;
}
let workers = image_fetch_concurrency().min(urls.len());
let next = std::sync::atomic::AtomicUsize::new(0);
std::thread::scope(|scope| {
for _ in 0..workers {
scope.spawn(|| loop {
let i = next.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
match urls.get(i) {
Some(url) => {
let _ = self.fetch_cached(url);
}
None => break,
}
});
}
});
}
}
pub(crate) struct MapImageResolver {
images: HashMap<String, PictureImage>,
}
impl MapImageResolver {
pub(crate) fn new(images: HashMap<String, PictureImage>) -> Self {
Self { images }
}
}
impl ImageResolver for MapImageResolver {
fn resolve(&self, src: &str) -> Option<PictureImage> {
self.images.get(src.trim()).cloned()
}
}
fn from_data_uri(uri: &str) -> Option<PictureImage> {
let rest = uri.strip_prefix("data:")?;
let (meta, payload) = rest.split_once(',')?;
let mime = meta
.split(';')
.next()
.filter(|m| !m.is_empty())
.unwrap_or("image/png");
let data = if meta.split(';').any(|t| t.eq_ignore_ascii_case("base64")) {
docling_core::base64::decode(payload)?
} else {
percent_decode(payload)
};
build_picture(mime, data)
}
#[cfg(feature = "fetch-images")]
fn is_blocked_ip(ip: std::net::IpAddr) -> bool {
use std::net::IpAddr;
match ip {
IpAddr::V4(v4) => {
v4.is_loopback()
|| v4.is_private()
|| v4.is_link_local()
|| v4.is_unspecified()
|| v4.is_broadcast()
|| v4.is_documentation()
|| (v4.octets()[0] == 100 && (v4.octets()[1] & 0xc0) == 64)
}
IpAddr::V6(v6) => {
v6.is_loopback()
|| v6.is_unspecified()
|| (v6.segments()[0] & 0xfe00) == 0xfc00
|| (v6.segments()[0] & 0xffc0) == 0xfe80
|| v6
.to_ipv4_mapped()
.is_some_and(|v4| is_blocked_ip(IpAddr::V4(v4)))
}
}
}
#[cfg(feature = "fetch-images")]
fn blocked_by_ssrf_guard(url: &str) -> bool {
use std::net::ToSocketAddrs;
let allow = std::env::var("DOCLING_RS_ALLOW_PRIVATE_IP_FETCH")
.map(|v| !v.is_empty() && v != "0" && !v.eq_ignore_ascii_case("false"))
.unwrap_or(false);
if allow {
return false;
}
let Ok(parsed) = url::Url::parse(url) else {
return true;
};
let Some(host) = parsed.host_str() else {
return true;
};
let port = parsed.port_or_known_default().unwrap_or(80);
match (host, port).to_socket_addrs() {
Ok(addrs) => addrs.map(|a| a.ip()).any(is_blocked_ip),
Err(_) => true,
}
}
#[cfg(feature = "fetch-images")]
fn fetch_remote(url: &str) -> Option<PictureImage> {
use std::time::Duration;
if blocked_by_ssrf_guard(url) {
return None;
}
let agent: ureq::Agent = ureq::Agent::config_builder()
.timeout_connect(Some(Duration::from_secs(5)))
.timeout_global(Some(Duration::from_secs(20)))
.max_redirects(3)
.build()
.into();
let mut resp = agent.get(url).call().ok()?;
let content_type = resp
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.map(|c| {
c.split(';')
.next()
.unwrap_or("")
.trim()
.to_ascii_lowercase()
});
let data = resp
.body_mut()
.with_config()
.limit(MAX_IMAGE_BYTES)
.read_to_vec()
.ok()?;
match content_type {
Some(mime) if mime.starts_with("image/") => build_picture(mime, data),
_ => {
let path = url.split(['?', '#']).next().unwrap_or(url);
super::ooxml::picture_image(path, data)
}
}
}
pub(crate) fn build_picture(mimetype: impl Into<String>, data: Vec<u8>) -> Option<PictureImage> {
if data.is_empty() {
return None;
}
let (width, height) = image::ImageReader::new(Cursor::new(&data))
.with_guessed_format()
.ok()?
.into_dimensions()
.ok()?;
Some(PictureImage {
mimetype: mimetype.into(),
width,
height,
data,
})
}
fn percent_decode(s: &str) -> Vec<u8> {
let b = s.as_bytes();
let mut out = Vec::with_capacity(b.len());
let mut i = 0;
while i < b.len() {
if b[i] == b'%' && i + 2 < b.len() {
if let (Some(h), Some(l)) = (hex(b[i + 1]), hex(b[i + 2])) {
out.push((h << 4) | l);
i += 3;
continue;
}
}
out.push(b[i]);
i += 1;
}
out
}
fn hex(c: u8) -> Option<u8> {
match c {
b'0'..=b'9' => Some(c - b'0'),
b'a'..=b'f' => Some(c - b'a' + 10),
b'A'..=b'F' => Some(c - b'A' + 10),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use docling_core::base64::encode;
const RED_PNG: &[u8] = &[
0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44,
0x52, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x02, 0x00, 0x00, 0x00, 0x90,
0x77, 0x53, 0xde, 0x00, 0x00, 0x00, 0x0c, 0x49, 0x44, 0x41, 0x54, 0x08, 0xd7, 0x63, 0xf8,
0xcf, 0xc0, 0x00, 0x00, 0x00, 0x03, 0x00, 0x01, 0x6e, 0x2c, 0xdc, 0x33, 0x00, 0x00, 0x00,
0x00, 0x49, 0x45, 0x4e, 0x44, 0xae, 0x42, 0x60, 0x82,
];
#[test]
fn decodes_base64_data_uri() {
let uri = format!("data:image/png;base64,{}", encode(RED_PNG));
let img = from_data_uri(&uri).expect("decodes");
assert_eq!(img.mimetype, "image/png");
assert_eq!((img.width, img.height), (1, 1));
assert_eq!(img.data, RED_PNG);
}
#[test]
fn rejects_garbage_data_uri() {
assert!(from_data_uri("data:image/png;base64,not-an-image").is_none());
assert!(from_data_uri("data:,").is_none());
}
#[test]
fn nofetch_resolves_nothing() {
assert!(NoFetch.resolve("data:image/png;base64,AAAA").is_none());
}
#[test]
fn map_resolver_returns_by_key() {
let img = from_data_uri(&format!("data:image/png;base64,{}", encode(RED_PNG))).unwrap();
let mut map = HashMap::new();
map.insert("images/x.png".to_string(), img.clone());
let r = MapImageResolver::new(map);
assert_eq!(r.resolve("images/x.png"), Some(img));
assert!(r.resolve("images/missing.png").is_none());
}
#[test]
fn fs_resolver_reads_absolute_files_but_not_relative_without_base() {
let p = std::env::temp_dir().join(format!("docling.rs_img_{}.png", std::process::id()));
std::fs::write(&p, RED_PNG).unwrap();
let r = FsImageResolver::new(None, None);
let img = r.resolve(p.to_str().unwrap()).expect("reads local file");
assert_eq!((img.width, img.height), (1, 1));
let _ = std::fs::remove_file(&p);
assert!(FsImageResolver::new(None, None)
.resolve("nope/relative.png")
.is_none());
assert!(r
.resolve(&format!("data:image/png;base64,{}", encode(RED_PNG)))
.is_some());
}
#[cfg(feature = "fetch-images")]
#[test]
fn resolves_relative_and_protocol_relative_against_base_url() {
let r = FsImageResolver::new(None, Some("https://ex.com/a/page.html".into()));
assert_eq!(
r.absolute_http_url("/img/x.png").as_deref(),
Some("https://ex.com/img/x.png")
);
assert_eq!(
r.absolute_http_url("y.png").as_deref(),
Some("https://ex.com/a/y.png")
);
assert_eq!(
r.absolute_http_url("//cdn.ex.com/z.png").as_deref(),
Some("https://cdn.ex.com/z.png")
);
assert_eq!(
r.absolute_http_url("https://other.com/w.png").as_deref(),
Some("https://other.com/w.png")
);
let no_base = FsImageResolver::new(None, None);
assert!(no_base.absolute_http_url("/img/x.png").is_none());
}
#[cfg(feature = "fetch-images")]
#[test]
fn concurrency_env_is_parsed_and_clamped() {
use super::image_fetch_concurrency;
std::env::set_var("DOCLING_RS_IMAGE_FETCH_CONCURRENCY", "7");
assert_eq!(image_fetch_concurrency(), 7);
std::env::set_var("DOCLING_RS_IMAGE_FETCH_CONCURRENCY", "0");
assert_eq!(image_fetch_concurrency(), 10, "0 → default, never zero");
std::env::set_var("DOCLING_RS_IMAGE_FETCH_CONCURRENCY", "9999");
assert_eq!(image_fetch_concurrency(), 64, "clamped to the ceiling");
std::env::set_var("DOCLING_RS_IMAGE_FETCH_CONCURRENCY", "junk");
assert_eq!(image_fetch_concurrency(), 10);
std::env::remove_var("DOCLING_RS_IMAGE_FETCH_CONCURRENCY");
}
#[cfg(feature = "fetch-images")]
#[test]
fn prefetch_fetches_each_url_once_and_warms_the_cache() {
use std::io::{Read, Write};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let hits = Arc::new(AtomicUsize::new(0));
let server_hits = Arc::clone(&hits);
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(mut stream) = stream else { break };
let mut buf = [0u8; 1024];
let _ = stream.read(&mut buf); server_hits.fetch_add(1, Ordering::Relaxed);
let header = format!(
"HTTP/1.1 200 OK\r\nContent-Type: image/png\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
RED_PNG.len()
);
let _ = stream.write_all(header.as_bytes());
let _ = stream.write_all(RED_PNG);
let _ = stream.flush();
}
});
std::env::set_var("DOCLING_RS_ALLOW_PRIVATE_IP_FETCH", "1");
let base = format!("http://{addr}/dir/page.html");
let r = FsImageResolver::new(None, Some(base));
r.prefetch(&[
"img1.png".to_string(),
"img2.png".to_string(),
"img1.png".to_string(),
]);
assert_eq!(
hits.load(Ordering::Relaxed),
2,
"each distinct URL fetched once"
);
let img = r.resolve("img1.png").expect("cached image resolves");
assert_eq!((img.width, img.height), (1, 1));
assert_eq!(
r.resolve("/dir/img2.png").map(|i| (i.width, i.height)),
Some((1, 1))
);
assert_eq!(
hits.load(Ordering::Relaxed),
2,
"resolve served from cache, no refetch"
);
std::env::remove_var("DOCLING_RS_ALLOW_PRIVATE_IP_FETCH");
}
}
#[cfg(not(feature = "fetch-images"))]
fn fetch_remote(_url: &str) -> Option<PictureImage> {
None
}