systemprompt_api/services/gateway/image_fetch/
mod.rs1mod guard;
20
21use base64::Engine as _;
22use base64::engine::general_purpose::STANDARD as BASE64;
23use systemprompt_models::net::{HTTP_CONNECT_TIMEOUT, trusted_http_hosts_from_env};
24
25use super::protocol::canonical::{CanonicalContent, CanonicalRequest, ImageSource};
26
27pub const MAX_IMAGE_BYTES: usize = 5 * 1024 * 1024;
30
31pub const ACCEPTED_MIME: [&str; 5] = [
32 "image/png",
33 "image/jpeg",
34 "image/webp",
35 "image/heic",
36 "image/heif",
37];
38
39const MAX_REDIRECTS: u8 = 3;
40const FETCH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10);
41
42#[derive(Debug, thiserror::Error)]
48#[error("image url {url} could not be inlined: {message}")]
49pub struct ImageFetchFailed {
50 pub url: String,
51 pub message: String,
52 pub caller_fault: bool,
53}
54
55#[derive(Debug, Clone)]
59pub struct ImageFetchPolicy {
60 pub timeout: std::time::Duration,
61 pub max_bytes: usize,
62 pub max_redirects: u8,
63 pub trusted_hosts: Vec<String>,
64}
65
66impl Default for ImageFetchPolicy {
67 fn default() -> Self {
68 Self {
69 timeout: FETCH_TIMEOUT,
70 max_bytes: MAX_IMAGE_BYTES,
71 max_redirects: MAX_REDIRECTS,
72 trusted_hosts: trusted_http_hosts_from_env(),
73 }
74 }
75}
76
77#[derive(Debug, Clone)]
79pub struct InlineImage {
80 pub media_type: String,
81 pub base64: String,
82}
83
84fn client() -> &'static reqwest::Client {
85 static CLIENT: std::sync::OnceLock<reqwest::Client> = std::sync::OnceLock::new();
86 CLIENT.get_or_init(|| {
87 reqwest::Client::builder()
88 .redirect(reqwest::redirect::Policy::none())
89 .connect_timeout(HTTP_CONNECT_TIMEOUT)
90 .build()
91 .unwrap_or_default()
92 })
93}
94
95pub async fn inline_url_images(
96 request: &mut CanonicalRequest,
97 policy: &ImageFetchPolicy,
98) -> Result<usize, ImageFetchFailed> {
99 let mut count = 0usize;
100 for message in &mut request.messages {
101 for content in &mut message.content {
102 let CanonicalContent::Image(ImageSource::Url { url, detail }) = content else {
103 continue;
104 };
105 let fetched = fetch(url, policy).await?;
106 *content = CanonicalContent::Image(ImageSource::Base64 {
107 media_type: fetched.media_type,
108 data: fetched.base64,
109 detail: *detail,
110 });
111 count += 1;
112 }
113 }
114 Ok(count)
115}
116
117pub async fn fetch(url: &str, policy: &ImageFetchPolicy) -> Result<InlineImage, ImageFetchFailed> {
118 let fail = |message: String, caller_fault: bool| ImageFetchFailed {
119 url: url.to_owned(),
120 message,
121 caller_fault,
122 };
123 tokio::time::timeout(policy.timeout, fetch_inner(url, policy))
124 .await
125 .map_or_else(
126 |_| Err(fail(format!("fetch exceeded {:?}", policy.timeout), false)),
127 |result| result.map_err(|(message, caller_fault)| fail(message, caller_fault)),
128 )
129}
130
131type FetchError = (String, bool);
132
133async fn fetch_inner(url: &str, policy: &ImageFetchPolicy) -> Result<InlineImage, FetchError> {
134 let mut next = guard::checked_url(url, &policy.trusted_hosts)
135 .await
136 .map_err(|e| (e, true))?;
137 for _ in 0..=policy.max_redirects {
138 let response = client()
139 .get(next.clone())
140 .send()
141 .await
142 .map_err(|e| (format!("request failed: {e}"), false))?;
143 if let Some(location) = redirect_target(&response) {
144 let joined = next
145 .join(&location)
146 .map_err(|e| (format!("invalid redirect target: {e}"), true))?;
147 next = guard::checked_url(joined.as_str(), &policy.trusted_hosts)
148 .await
149 .map_err(|e| (format!("redirect rejected: {e}"), true))?;
150 continue;
151 }
152 return read_image(response, policy).await;
153 }
154 Err((
155 format!("more than {} redirects", policy.max_redirects),
156 true,
157 ))
158}
159
160fn redirect_target(response: &reqwest::Response) -> Option<String> {
161 if !response.status().is_redirection() {
162 return None;
163 }
164 response
165 .headers()
166 .get(reqwest::header::LOCATION)
167 .and_then(|v| v.to_str().ok())
168 .map(ToOwned::to_owned)
169}
170
171async fn read_image(
172 mut response: reqwest::Response,
173 policy: &ImageFetchPolicy,
174) -> Result<InlineImage, FetchError> {
175 let status = response.status();
176 if !status.is_success() {
177 return Err((format!("host returned {status}"), true));
178 }
179 let media_type = declared_mime(&response)?;
180 let mut body: Vec<u8> = Vec::new();
181 while let Some(chunk) = response
182 .chunk()
183 .await
184 .map_err(|e| (format!("read failed: {e}"), false))?
185 {
186 if body.len() + chunk.len() > policy.max_bytes {
187 return Err((format!("larger than {} bytes", policy.max_bytes), true));
188 }
189 body.extend_from_slice(&chunk);
190 }
191 if body.is_empty() {
192 return Err(("empty response body".to_owned(), true));
193 }
194 Ok(InlineImage {
195 media_type,
196 base64: BASE64.encode(&body),
197 })
198}
199
200fn declared_mime(response: &reqwest::Response) -> Result<String, FetchError> {
201 let raw = response
202 .headers()
203 .get(reqwest::header::CONTENT_TYPE)
204 .and_then(|v| v.to_str().ok())
205 .ok_or_else(|| ("no content-type".to_owned(), true))?;
206 let mime = raw
207 .split(';')
208 .next()
209 .unwrap_or_default()
210 .trim()
211 .to_ascii_lowercase();
212 if ACCEPTED_MIME.contains(&mime.as_str()) {
213 return Ok(mime);
214 }
215 Err((
216 format!("content-type {mime} is not an inlineable image"),
217 true,
218 ))
219}