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