fix(providers): bound local image reads (#11114)

This commit is contained in:
Jasper
2026-08-12 15:29:06 -05:00
committed by GitHub
parent 4e28d45864
commit e1e46b1568
+68 -23
View File
@@ -7,6 +7,8 @@ use serde_json::{json, Value};
use crate::errors::ProviderError;
const MAX_IMAGE_BYTES: u64 = 20 * 1024 * 1024;
#[derive(Debug, Copy, Clone, Serialize, Deserialize)]
pub enum ImageFormat {
OpenAi,
@@ -183,37 +185,31 @@ fn is_image_file(path: &Path) -> bool {
if let Ok(mut file) = std::fs::File::open(path) {
let mut buffer = [0u8; 8]; // Large enough for most image magic numbers
if file.read(&mut buffer).is_ok() {
// Check magic numbers for common image formats
return match &buffer[0..4] {
// PNG: 89 50 4E 47
[0x89, 0x50, 0x4E, 0x47] => true,
// JPEG: FF D8 FF
[0xFF, 0xD8, 0xFF, _] => true,
// GIF: 47 49 46 38
[0x47, 0x49, 0x46, 0x38] => true,
_ => false,
};
return has_image_magic(&buffer);
}
}
false
}
fn has_image_magic(bytes: &[u8]) -> bool {
matches!(
bytes.get(..4),
Some([0x89, 0x50, 0x4E, 0x47])
| Some([0xFF, 0xD8, 0xFF, _])
| Some([0x47, 0x49, 0x46, 0x38])
)
}
fn read_bounded(reader: impl std::io::Read, max_bytes: u64) -> std::io::Result<Vec<u8>> {
let mut bytes = Vec::new();
reader.take(max_bytes + 1).read_to_end(&mut bytes)?;
Ok(bytes)
}
/// Convert a local image file to base64 encoded ImageContent
pub fn load_image_file(path: &str) -> Result<ImageContent, ProviderError> {
let path = Path::new(path);
// Verify it's an image before proceeding
if !is_image_file(path) {
return Err(ProviderError::RequestFailed(
"File is not a valid image".to_string(),
));
}
// Read the file
let bytes = std::fs::read(path)
.map_err(|e| ProviderError::RequestFailed(format!("Failed to read image file: {}", e)))?;
// Detect mime type from extension
let mime_type = match path.extension().and_then(|e| e.to_str()) {
Some(ext) => match ext.to_lowercase().as_str() {
"png" => "image/png",
@@ -231,7 +227,33 @@ pub fn load_image_file(path: &str) -> Result<ImageContent, ProviderError> {
}
};
// Convert to base64
let file = std::fs::File::open(path)
.map_err(|e| ProviderError::RequestFailed(format!("Failed to read image file: {e}")))?;
let file_size = file
.metadata()
.map_err(|e| ProviderError::RequestFailed(format!("Failed to read image file: {e}")))?
.len();
if file_size > MAX_IMAGE_BYTES {
return Err(ProviderError::RequestFailed(format!(
"Image file exceeds the {} MiB limit",
MAX_IMAGE_BYTES / (1024 * 1024)
)));
}
let bytes = read_bounded(file, MAX_IMAGE_BYTES)
.map_err(|e| ProviderError::RequestFailed(format!("Failed to read image file: {e}")))?;
if bytes.len() as u64 > MAX_IMAGE_BYTES {
return Err(ProviderError::RequestFailed(format!(
"Image file exceeds the {} MiB limit",
MAX_IMAGE_BYTES / (1024 * 1024)
)));
}
if !has_image_magic(&bytes) {
return Err(ProviderError::RequestFailed(
"File is not a valid image".to_string(),
));
}
let data = base64::prelude::BASE64_STANDARD.encode(&bytes);
Ok(ImageContent::new(data, mime_type))
@@ -240,6 +262,7 @@ pub fn load_image_file(path: &str) -> Result<ImageContent, ProviderError> {
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write as _;
use tempfile;
#[test]
@@ -476,6 +499,10 @@ mod tests {
assert!(result.is_ok());
let image = result.unwrap();
assert_eq!(image.mime_type, "image/png");
assert_eq!(
base64::prelude::BASE64_STANDARD.decode(image.data).unwrap(),
png_data
);
// Test loading fake PNG file
let result = load_image_file(fake_png_path_str);
@@ -504,4 +531,22 @@ mod tests {
.to_string()
.contains("Unsupported image format"));
}
#[test]
fn bounded_reader_accepts_limit_and_detects_extra_byte() {
assert_eq!(read_bounded(&b"12345678"[..], 8).unwrap().len(), 8);
assert_eq!(read_bounded(&b"123456789"[..], 8).unwrap().len(), 9);
}
#[test]
fn load_image_file_rejects_oversized_sparse_file() {
let temp_dir = tempfile::tempdir().unwrap();
let png_path = temp_dir.path().join("oversized.png");
let mut file = std::fs::File::create(&png_path).unwrap();
file.write_all(&[0x89, 0x50, 0x4E, 0x47]).unwrap();
file.set_len(MAX_IMAGE_BYTES + 1).unwrap();
let error = load_image_file(png_path.to_str().unwrap()).unwrap_err();
assert!(error.to_string().contains("exceeds the 20 MiB limit"));
}
}