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}