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::{GuardedConnectError, 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
84pub async fn inline_url_images(
85 request: &mut CanonicalRequest,
86 policy: &ImageFetchPolicy,
87) -> Result<usize, ImageFetchFailed> {
88 let mut count = 0usize;
89 for message in &mut request.messages {
90 for content in &mut message.content {
91 let CanonicalContent::Image(ImageSource::Url { url, detail }) = content else {
92 continue;
93 };
94 let fetched = fetch(url, policy).await?;
95 *content = CanonicalContent::Image(ImageSource::Base64 {
96 media_type: fetched.media_type,
97 data: fetched.base64,
98 detail: *detail,
99 });
100 count += 1;
101 }
102 }
103 Ok(count)
104}
105
106pub async fn fetch(url: &str, policy: &ImageFetchPolicy) -> Result<InlineImage, ImageFetchFailed> {
107 let fail = |message: String, caller_fault: bool| ImageFetchFailed {
108 url: url.to_owned(),
109 message,
110 caller_fault,
111 };
112 tokio::time::timeout(policy.timeout, fetch_inner(url, policy))
113 .await
114 .map_or_else(
115 |_| Err(fail(format!("fetch exceeded {:?}", policy.timeout), false)),
116 |result| result.map_err(|(message, caller_fault)| fail(message, caller_fault)),
117 )
118}
119
120type FetchError = (String, bool);
121
122async fn fetch_inner(url: &str, policy: &ImageFetchPolicy) -> Result<InlineImage, FetchError> {
123 let checked = guard::checked_url(url, &policy.trusted_hosts).map_err(|e| (e, true))?;
124 let client = guard::client(policy).map_err(|e| (e, false))?;
125 let response = client
126 .get(checked)
127 .send()
128 .await
129 .map_err(|e| describe_send_error(&e, policy))?;
130 read_image(response, policy).await
131}
132
133fn describe_send_error(error: &reqwest::Error, policy: &ImageFetchPolicy) -> FetchError {
136 if error.is_timeout() {
137 return (format!("fetch exceeded {:?}", policy.timeout), false);
138 }
139 let mut source = std::error::Error::source(error);
140 while let Some(inner) = source {
141 if let Some(guarded) = inner.downcast_ref::<GuardedConnectError>() {
142 return match guarded {
143 GuardedConnectError::RedirectRefused { .. }
144 | GuardedConnectError::TooManyRedirects(_) => {
145 (format!("redirect rejected: {guarded}"), true)
146 },
147 GuardedConnectError::Unresolvable(_)
148 | GuardedConnectError::BlockedAddress { .. } => {
149 (format!("host rejected: {guarded}"), true)
150 },
151 };
152 }
153 source = inner.source();
154 }
155 (format!("request failed: {error}"), error.is_redirect())
156}
157
158async fn read_image(
159 mut response: reqwest::Response,
160 policy: &ImageFetchPolicy,
161) -> Result<InlineImage, FetchError> {
162 let status = response.status();
163 if !status.is_success() {
164 return Err((format!("host returned {status}"), true));
165 }
166 let media_type = declared_mime(&response)?;
167 let mut body: Vec<u8> = Vec::new();
168 while let Some(chunk) = response
169 .chunk()
170 .await
171 .map_err(|e| (format!("read failed: {e}"), false))?
172 {
173 if body.len() + chunk.len() > policy.max_bytes {
174 return Err((format!("larger than {} bytes", policy.max_bytes), true));
175 }
176 body.extend_from_slice(&chunk);
177 }
178 if body.is_empty() {
179 return Err(("empty response body".to_owned(), true));
180 }
181 Ok(InlineImage {
182 media_type,
183 base64: BASE64.encode(&body),
184 })
185}
186
187fn declared_mime(response: &reqwest::Response) -> Result<String, FetchError> {
188 let raw = response
189 .headers()
190 .get(reqwest::header::CONTENT_TYPE)
191 .and_then(|v| v.to_str().ok())
192 .ok_or_else(|| ("no content-type".to_owned(), true))?;
193 let mime = raw
194 .split(';')
195 .next()
196 .unwrap_or_default()
197 .trim()
198 .to_ascii_lowercase();
199 if ACCEPTED_MIME.contains(&mime.as_str()) {
200 return Ok(mime);
201 }
202 Err((
203 format!("content-type {mime} is not an inlineable image"),
204 true,
205 ))
206}