use crate::util::UnwrapPoison;
use base64::{
Engine as _,
engine::general_purpose::{STANDARD, URL_SAFE, URL_SAFE_NO_PAD},
};
use std::borrow::Cow;
use std::collections::HashMap;
use std::io::Cursor;
use std::path::{Path, PathBuf};
use std::sync::{LazyLock, Mutex};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum MediaTarget {
LocalImage,
RemoteUrl,
Invalid,
}
pub(crate) const MAX_DATA_URI_ENCODED_BYTES: usize = 20 * 1024 * 1024;
pub(crate) const RASTER_DECODE_MAX_ALLOC_BYTES: u64 = 256 * 1024 * 1024;
pub(crate) const RASTER_DECODE_MAX_DIMENSION_PX: u32 = 16384;
#[must_use]
pub(crate) fn raster_decode_limits() -> image::Limits {
let mut limits = image::Limits::default();
limits.max_alloc = Some(RASTER_DECODE_MAX_ALLOC_BYTES);
limits.max_image_width = Some(RASTER_DECODE_MAX_DIMENSION_PX);
limits.max_image_height = Some(RASTER_DECODE_MAX_DIMENSION_PX);
limits
}
#[must_use]
pub(crate) fn classify_media_image_target(target: &str) -> MediaTarget {
let trimmed = target.trim();
if is_valid_remote_url(trimmed) {
MediaTarget::RemoteUrl
} else if is_decodable_local_raster(trimmed) {
MediaTarget::LocalImage
} else {
MediaTarget::Invalid
}
}
#[must_use]
pub(crate) fn is_valid_remote_url(target: &str) -> bool {
let Some(rest) = target
.strip_prefix("https://")
.or_else(|| target.strip_prefix("http://"))
else {
return false;
};
if rest.chars().any(char::is_whitespace) {
return false;
}
let authority = rest.split(['/', '?', '#']).next().unwrap_or("");
if authority.is_empty() || authority.starts_with(':') || authority.contains('@') {
return false;
}
let Some((host, port)) = split_authority_host_port(authority) else {
return false;
};
if !port.is_empty()
&& (!port.bytes().all(|b| b.is_ascii_digit()) || !port.parse::<u16>().is_ok_and(|x| x > 0))
{
return false;
}
if host.starts_with('[') {
host.ends_with(']')
&& host.len() > 2
&& host.matches('[').count() == 1
&& host.matches(']').count() == 1
} else {
!host.is_empty() && !host.contains(':') && !host.contains('[')
}
}
fn split_authority_host_port(authority: &str) -> Option<(&str, &str)> {
if let Some(close) = authority.find(']') {
if !authority.starts_with('[') {
return None;
}
let after = &authority[close + 1..];
if after.is_empty() {
return Some((&authority[..=close], ""));
}
let port = after.strip_prefix(':')?;
if port.is_empty() {
return None; }
Some((&authority[..=close], port))
} else {
match authority.rsplit_once(':') {
Some((host, port)) if !port.is_empty() => Some((host, port)),
Some(_) => None, None => Some((authority, "")),
}
}
}
#[must_use]
pub(crate) fn data_uri_base64_payload(uri: &str) -> Option<&str> {
let rest = uri.strip_prefix("data:")?;
let payload_start = rest.rfind(";base64,")? + ";base64,".len();
let payload = &rest[payload_start..];
(payload.len() <= MAX_DATA_URI_ENCODED_BYTES).then_some(payload)
}
fn declared_data_uri_format(target: &str) -> Option<image::ImageFormat> {
let rest = target.strip_prefix("data:image/")?;
let subtype = rest.split(';').next().unwrap_or("").trim();
let declared = image::ImageFormat::from_extension(subtype)?;
crate::util::image_format_native_label(declared)
.is_some()
.then_some(declared)
}
fn parse_native_data_uri(uri: &str) -> Option<(image::ImageFormat, Vec<u8>)> {
let declared = declared_data_uri_format(uri)?;
let payload = data_uri_base64_payload(uri)?;
let bytes = decode_base64_payload(payload)?;
if image::guess_format(&bytes).ok()? != declared {
return None;
}
Some((declared, bytes))
}
#[must_use]
pub(crate) fn decode_native_data_uri(
uri: &str,
limits: image::Limits,
) -> Option<(u32, u32, Vec<u8>)> {
let (_, bytes) = parse_native_data_uri(uri)?;
decode_raster_bytes(&bytes, limits)
}
#[must_use]
pub(crate) fn is_native_data_uri(target: &str) -> bool {
parse_native_data_uri(target)
.is_some_and(|(_, bytes)| decode_raster(&bytes, raster_decode_limits()).is_some())
}
#[must_use]
fn decode_base64_payload(s: &str) -> Option<Vec<u8>> {
let compact: Cow<'_, [u8]> = if s.as_bytes().iter().any(u8::is_ascii_whitespace) {
Cow::Owned(
s.bytes()
.filter(|b| !b.is_ascii_whitespace())
.collect::<Vec<u8>>(),
)
} else {
Cow::Borrowed(s.as_bytes())
};
if compact.is_empty() {
return None;
}
STANDARD
.decode(compact.as_ref())
.ok()
.or_else(|| URL_SAFE.decode(compact.as_ref()).ok())
.or_else(|| URL_SAFE_NO_PAD.decode(compact.as_ref()).ok())
}
#[must_use]
fn decode_raster(bytes: &[u8], limits: image::Limits) -> Option<image::DynamicImage> {
let format = image::guess_format(bytes).ok()?;
crate::util::image_format_native_label(format)?;
let mut reader = image::ImageReader::with_format(Cursor::new(bytes), format);
reader.limits(limits);
reader.decode().ok()
}
#[must_use]
pub(crate) fn decode_raster_bytes(
bytes: &[u8],
limits: image::Limits,
) -> Option<(u32, u32, Vec<u8>)> {
let img = decode_raster(bytes, limits)?;
let rgba = img.to_rgba8();
Some((rgba.width(), rgba.height(), rgba.into_raw()))
}
static LOCAL_RASTER_DECODE_CACHE: LazyLock<Mutex<HashMap<(PathBuf, std::time::SystemTime), bool>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
const LOCAL_RASTER_CACHE_MAX_ENTRIES: usize = 4096;
#[must_use]
fn is_decodable_local_raster(target: &str) -> bool {
let path = Path::new(target);
if !path.is_absolute() {
return false;
}
let Ok(meta) = std::fs::metadata(path) else {
return false;
};
if !meta.is_file() || meta.len() > crate::util::INBOUND_IMAGE_MAX_INPUT_BYTES {
return false;
}
let Ok(mtime) = meta.modified() else {
return false;
};
let canonical = std::fs::canonicalize(path).unwrap_or_else(|_| path.to_path_buf());
let key = (canonical.clone(), mtime);
{
let cache = LOCAL_RASTER_DECODE_CACHE.lock().unwrap_poison();
if let Some(&result) = cache.get(&key) {
return result;
}
}
let result = is_decodable_raster_file(&canonical);
let mut cache = LOCAL_RASTER_DECODE_CACHE.lock().unwrap_poison();
if cache.len() >= LOCAL_RASTER_CACHE_MAX_ENTRIES {
cache.clear();
}
cache.insert(key, result);
result
}
#[must_use]
fn is_decodable_raster_file(path: &Path) -> bool {
std::fs::read(path).is_ok_and(|bytes| decode_raster(&bytes, raster_decode_limits()).is_some())
}
#[cfg(test)]
mod tests {
use super::*;
use base64::engine::general_purpose::STANDARD;
fn write_temp_png(tag: &str) -> std::path::PathBuf {
use std::io::Write;
let path = std::env::temp_dir().join(format!(
"mahbot_media_target_{tag}_{}.png",
std::process::id()
));
let img = image::RgbaImage::from_pixel(1, 1, image::Rgba([255, 0, 0, 255]));
let mut buf = Vec::new();
img.write_to(&mut std::io::Cursor::new(&mut buf), image::ImageFormat::Png)
.expect("test PNG must encode");
std::fs::File::create(&path)
.and_then(|mut f| f.write_all(&buf))
.expect("test PNG must write");
path
}
#[test]
fn remote_url_valid_and_malformed() {
assert_eq!(
classify_media_image_target("https://example.com/b.jpg"),
MediaTarget::RemoteUrl
);
assert_eq!(
classify_media_image_target("http://127.0.0.1:8080/img.png"),
MediaTarget::RemoteUrl
);
assert_eq!(classify_media_image_target("http://"), MediaTarget::Invalid);
assert_eq!(
classify_media_image_target("https://exa mple.com/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https:// example.com/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("http://:8080/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https://example.com:99999/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https://example.com:/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https://[::1/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https://user:pass@host/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https://user@host/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("http://[::1]:8080/img.png"),
MediaTarget::RemoteUrl
);
assert_eq!(
classify_media_image_target("https://host:8080:9999/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("http://[::1]:80:90/img.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("http://[]/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("http://[::1]foo:80/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https://host[::1]:80/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("https://[::1]/img.png"),
MediaTarget::RemoteUrl
);
assert_eq!(
classify_media_image_target("https://example.com/foo bar.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("http://host[/x.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("http://host[foo:80"),
MediaTarget::Invalid
);
}
#[test]
fn data_uri_native_and_fake() {
let img = image::RgbaImage::from_pixel(1, 1, image::Rgba([255, 0, 0, 255]));
let mut buf = Vec::new();
img.write_to(&mut std::io::Cursor::new(&mut buf), image::ImageFormat::Png)
.expect("test PNG must encode");
let uri = format!("data:image/png;base64,{}", STANDARD.encode(&buf));
assert_eq!(classify_media_image_target(&uri), MediaTarget::Invalid);
assert!(is_native_data_uri(&uri));
assert!(!is_native_data_uri("data:image/png;base64,abcd"));
let truncated = {
let payload = uri
.strip_prefix("data:image/png;base64,")
.expect("tiny png uri");
let bytes = STANDARD
.decode(payload.as_bytes())
.expect("tiny png base64");
let t = &bytes[..bytes.len().min(24)];
format!("data:image/png;base64,{}", STANDARD.encode(t))
};
assert!(!is_native_data_uri(&truncated));
assert!(!is_native_data_uri("data:image/gif;base64,R0lGOD"));
let mut jbuf = Vec::new();
image::RgbImage::from_pixel(1, 1, image::Rgb([255, 0, 0]))
.write_to(
&mut std::io::Cursor::new(&mut jbuf),
image::ImageFormat::Jpeg,
)
.expect("test JPEG must encode");
let mismatched = format!("data:image/png;base64,{}", STANDARD.encode(&jbuf));
assert!(!is_native_data_uri(&mismatched));
}
#[test]
fn data_uri_base64_payload_extracts_and_caps() {
let cases: &[(&str, Option<&str>)] = &[
("data:image/png;base64,AAAA", Some("AAAA")),
("data:image/png;charset=utf-8;base64,BBBB", Some("BBBB")),
("/tmp/photo.png", None),
("https://example.com/img.png", None),
("data:image/png,raw-bytes", None),
("data:image/png;base64,", Some("")),
];
for (uri, expected) in cases {
assert_eq!(data_uri_base64_payload(uri), *expected, "case: {uri}");
}
let over = format!(
"data:image/png;base64,{}",
"A".repeat(MAX_DATA_URI_ENCODED_BYTES + 1)
);
assert_eq!(data_uri_base64_payload(&over), None);
}
#[test]
fn local_raster_and_non_image() {
let png = write_temp_png("valid");
let text = std::env::temp_dir().join(format!(
"mahbot_media_target_{}_notimage.txt",
std::process::id()
));
std::fs::write(&text, b"top secret").unwrap();
assert_eq!(
classify_media_image_target(png.to_str().unwrap()),
MediaTarget::LocalImage
);
assert_eq!(
classify_media_image_target(text.to_str().unwrap()),
MediaTarget::Invalid
);
let _ = std::fs::remove_file(&png);
let _ = std::fs::remove_file(&text);
}
#[test]
fn placeholders_and_relative_are_invalid() {
assert_eq!(classify_media_image_target(""), MediaTarget::Invalid);
assert_eq!(classify_media_image_target(" "), MediaTarget::Invalid);
assert_eq!(classify_media_image_target("..."), MediaTarget::Invalid);
assert_eq!(
classify_media_image_target("photo.png"),
MediaTarget::Invalid
);
assert_eq!(
classify_media_image_target("/tmp/definitely_missing_mahbot.png"),
MediaTarget::Invalid
);
}
#[test]
fn truncated_local_raster_is_invalid() {
let png = write_temp_png("trunc");
let bytes = std::fs::read(&png).unwrap();
let truncated = &bytes[..bytes.len().min(24)];
let corrupt = std::env::temp_dir().join(format!(
"mahbot_media_target_{}_trunc.png",
std::process::id()
));
std::fs::write(&corrupt, truncated).unwrap();
assert_eq!(
classify_media_image_target(corrupt.to_str().unwrap()),
MediaTarget::Invalid
);
let _ = std::fs::remove_file(&png);
let _ = std::fs::remove_file(&corrupt);
}
}