Skip to main content

iris_reader/
image.rs

1use std::{
2    collections::{HashMap, HashSet},
3    fs,
4    path::{Path, PathBuf},
5};
6
7use ::image::{DynamicImage, RgbaImage};
8use ratatui::layout::Size;
9use ratatui_image::{Resize, picker::Picker, protocol::Protocol};
10use resvg::{tiny_skia, usvg};
11
12#[derive(Debug, Clone, PartialEq, Eq, Hash)]
13enum ImageSource {
14    Local(PathBuf),
15    Remote(String),
16}
17
18#[derive(Debug, Clone, PartialEq, Eq, Hash)]
19struct CacheKey {
20    source: ImageSource,
21    width: u16,
22    height: u16,
23}
24
25#[derive(Default)]
26pub struct ImageManager {
27    picker: Option<Picker>,
28    decoded: HashMap<ImageSource, DynamicImage>,
29    protocols: HashMap<CacheKey, Protocol>,
30    failed_sources: HashSet<ImageSource>,
31}
32
33impl ImageManager {
34    pub fn initialize(&mut self) {
35        if self.picker.is_none() {
36            self.picker = Some(Picker::from_query_stdio().unwrap_or_else(|_| Picker::halfblocks()));
37        }
38    }
39
40    pub fn layout_size(
41        &mut self,
42        source: &str,
43        base_dir: Option<&Path>,
44        width: u16,
45        max_height: u16,
46    ) -> Option<Size> {
47        let source = resolve_source(source, base_dir)?;
48        self.ensure_decoded(&source)?;
49
50        let image = self.decoded.get(&source)?;
51        let picker = self.picker.as_ref()?;
52        let available = Size::new(width.max(1), max_height.max(1));
53
54        Some(Resize::Fit(None).size_for(image, picker.font_size(), available))
55    }
56
57    pub fn protocol(
58        &mut self,
59        source: &str,
60        base_dir: Option<&Path>,
61        width: u16,
62        height: u16,
63    ) -> Option<&Protocol> {
64        let source = resolve_source(source, base_dir)?;
65        if self.failed_sources.contains(&source) {
66            return None;
67        }
68
69        let key = CacheKey {
70            source: source.clone(),
71            width: width.max(1),
72            height: height.max(1),
73        };
74
75        if !self.protocols.contains_key(&key) {
76            self.ensure_decoded(&source)?;
77
78            let image = self.decoded.get(&source)?.clone();
79            let picker = self.picker.as_ref()?;
80            let protocol = picker
81                .new_protocol(image, Size::new(key.width, key.height), Resize::Fit(None))
82                .ok()?;
83
84            self.protocols.insert(key.clone(), protocol);
85        }
86
87        self.protocols.get(&key)
88    }
89
90    fn ensure_decoded(&mut self, source: &ImageSource) -> Option<()> {
91        if self.failed_sources.contains(source) {
92            return None;
93        }
94
95        if !self.decoded.contains_key(source) {
96            let Some(image) = load_image(source) else {
97                self.failed_sources.insert(source.clone());
98                return None;
99            };
100
101            self.decoded.insert(source.clone(), image);
102        }
103
104        Some(())
105    }
106}
107
108fn load_image(source: &ImageSource) -> Option<DynamicImage> {
109    match source {
110        ImageSource::Local(path) => {
111            let bytes = fs::read(path).ok()?;
112            decode_image(&bytes, Some(path), looks_like_svg_path(path))
113        }
114        ImageSource::Remote(url) => {
115            let mut response = ureq::get(url.as_str()).call().ok()?;
116            let bytes = response.body_mut().read_to_vec().ok()?;
117            decode_image(&bytes, None, looks_like_svg_url(url))
118        }
119    }
120}
121
122fn decode_image(bytes: &[u8], source_path: Option<&Path>, svg_hint: bool) -> Option<DynamicImage> {
123    if svg_hint || looks_like_svg_bytes(bytes) {
124        return rasterize_svg(bytes, source_path);
125    }
126
127    ::image::load_from_memory(bytes).ok()
128}
129
130fn rasterize_svg(bytes: &[u8], source_path: Option<&Path>) -> Option<DynamicImage> {
131    let resources_dir = source_path.and_then(Path::parent).map(Path::to_path_buf);
132
133    let mut options = usvg::Options {
134        resources_dir,
135        ..usvg::Options::default()
136    };
137    options.fontdb_mut().load_system_fonts();
138
139    let tree = usvg::Tree::from_data(bytes, &options).ok()?;
140    let size = tree.size().to_int_size();
141    let mut pixmap = tiny_skia::Pixmap::new(size.width(), size.height())?;
142
143    resvg::render(&tree, tiny_skia::Transform::default(), &mut pixmap.as_mut());
144
145    let pixels = pixmap.take_demultiplied();
146    let image = RgbaImage::from_raw(size.width(), size.height(), pixels)?;
147
148    Some(DynamicImage::ImageRgba8(image))
149}
150
151fn looks_like_svg_path(path: &Path) -> bool {
152    path.extension()
153        .and_then(|extension| extension.to_str())
154        .is_some_and(|extension| extension.eq_ignore_ascii_case("svg"))
155}
156
157fn looks_like_svg_url(url: &str) -> bool {
158    let path = url.split(['?', '#']).next().unwrap_or(url);
159
160    path.rsplit_once('.')
161        .is_some_and(|(_, extension)| extension.eq_ignore_ascii_case("svg"))
162}
163
164fn looks_like_svg_bytes(bytes: &[u8]) -> bool {
165    let prefix = String::from_utf8_lossy(&bytes[..bytes.len().min(4096)]);
166    let trimmed = prefix.trim_start_matches('\u{feff}').trim_start();
167
168    trimmed.starts_with("<svg")
169        || trimmed.starts_with("<?xml") && trimmed.contains("<svg")
170        || trimmed.starts_with("<!--") && trimmed.contains("<svg")
171}
172
173fn resolve_source(source: &str, base_dir: Option<&Path>) -> Option<ImageSource> {
174    if source.starts_with("http://") || source.starts_with("https://") {
175        return Some(ImageSource::Remote(source.to_string()));
176    }
177
178    if source.starts_with("data:") {
179        return None;
180    }
181
182    let source = source.strip_prefix("file://").unwrap_or(source);
183    let path = Path::new(source);
184    let path = if path.is_absolute() {
185        path.to_path_buf()
186    } else if let Some(base_dir) = base_dir {
187        base_dir.join(path)
188    } else {
189        path.to_path_buf()
190    };
191
192    Some(ImageSource::Local(path))
193}
194
195#[cfg(test)]
196mod tests {
197    use super::*;
198
199    #[test]
200    fn resolves_relative_images_from_document_directory() {
201        let base = Path::new("/tmp/iris-doc");
202        assert_eq!(
203            resolve_source("images/example.png", Some(base)),
204            Some(ImageSource::Local(base.join("images/example.png")))
205        );
206    }
207
208    #[test]
209    fn keeps_remote_images_remote() {
210        assert_eq!(
211            resolve_source("https://example.com/image.png", None),
212            Some(ImageSource::Remote(
213                "https://example.com/image.png".to_string()
214            ))
215        );
216    }
217
218    #[test]
219    fn ignores_data_urls_for_now() {
220        assert!(resolve_source("data:image/png;base64,AAAA", None).is_none());
221    }
222
223    #[test]
224    fn decodes_supported_local_image() {
225        let source = ImageSource::Local(PathBuf::from("tests/fixtures/assets/iris.png"));
226        let image = load_image(&source).expect("fixture image should decode");
227        assert_eq!(image.width(), 1);
228        assert_eq!(image.height(), 1);
229    }
230
231    #[test]
232    fn detects_svg_content_without_svg_extension() {
233        let svg = br#"<svg xmlns="http://www.w3.org/2000/svg" width="20" height="10"></svg>"#;
234        assert!(looks_like_svg_bytes(svg));
235    }
236
237    #[test]
238    fn rasterizes_svg_content() {
239        let svg = br##"
240            <svg xmlns="http://www.w3.org/2000/svg" width="20" height="10">
241                <rect width="20" height="10" fill="#ff7500" />
242            </svg>
243        "##;
244
245        let image = decode_image(svg, None, true).expect("SVG should rasterize");
246        assert_eq!(image.width(), 20);
247        assert_eq!(image.height(), 10);
248    }
249}