1use std::sync::Arc;
22use std::time::{Duration, Instant};
23
24use async_trait::async_trait;
25use base64::Engine;
26use serde::Deserialize;
27
28use crate::image_gen::{
29 GeneratedImage, ImageGenError, ImageGenProvider, ImageGenRequest, ImageGenResult, ImageSink,
30 ImageSize,
31};
32
33pub const FAL_DEFAULT_BASE: &str = "https://queue.fal.run";
35const POLL_INTERVAL: Duration = Duration::from_millis(3000);
37const MAX_WAIT: Duration = Duration::from_millis(175_000);
39
40pub struct FalImageProvider {
42 client: reqwest::Client,
43 base_url: String,
44 api_key: String,
45 store: Arc<dyn ImageSink>,
46}
47
48impl FalImageProvider {
49 pub fn new(
52 base_url: String,
53 api_key: String,
54 store: Arc<dyn ImageSink>,
55 ) -> Result<Self, ImageGenError> {
56 let client = reqwest::Client::builder()
57 .timeout(Duration::from_secs(30))
58 .connect_timeout(Duration::from_secs(10))
59 .build()
60 .map_err(|e| ImageGenError::Transport(e.to_string()))?;
61 Ok(Self {
62 client,
63 base_url,
64 api_key,
65 store,
66 })
67 }
68
69 fn auth(&self) -> String {
70 format!("Key {}", self.api_key)
71 }
72
73 fn submit_url(&self, model: &str) -> String {
76 format!(
77 "{}/{}",
78 self.base_url.trim_end_matches('/'),
79 model.trim_start_matches('/')
80 )
81 }
82
83 async fn authed_get<T: for<'de> serde::Deserialize<'de>>(
84 &self,
85 url: &str,
86 ) -> Result<T, ImageGenError> {
87 let resp = self
88 .client
89 .get(url)
90 .header(reqwest::header::AUTHORIZATION, self.auth())
91 .send()
92 .await
93 .map_err(|e| ImageGenError::Transport(e.to_string()))?;
94 let status = resp.status();
95 if !status.is_success() {
96 let body = resp.text().await.unwrap_or_default();
97 return Err(ImageGenError::Http {
98 status: status.as_u16(),
99 body,
100 });
101 }
102 resp.json::<T>()
103 .await
104 .map_err(|e| ImageGenError::BadResponse(format!("invalid JSON: {e}")))
105 }
106}
107
108#[async_trait]
109impl ImageGenProvider for FalImageProvider {
110 fn name(&self) -> &'static str {
111 "fal"
112 }
113
114 async fn generate(&self, req: &ImageGenRequest) -> Result<ImageGenResult, ImageGenError> {
115 let model = req.model.as_deref().ok_or(ImageGenError::MissingModel)?;
116
117 let submit_resp = self
119 .client
120 .post(self.submit_url(model))
121 .header(reqwest::header::AUTHORIZATION, self.auth())
122 .json(&build_input(req))
123 .send()
124 .await
125 .map_err(|e| ImageGenError::Transport(e.to_string()))?;
126 let status = submit_resp.status();
127 if !status.is_success() {
128 let body = submit_resp.text().await.unwrap_or_default();
129 return Err(ImageGenError::Http {
130 status: status.as_u16(),
131 body,
132 });
133 }
134 let submit: FalSubmit = submit_resp
135 .json()
136 .await
137 .map_err(|e| ImageGenError::BadResponse(format!("submit parse: {e}")))?;
138 let status_url = submit
139 .status_url
140 .ok_or_else(|| ImageGenError::BadResponse("submit returned no status_url".into()))?;
141 let response_url = submit
142 .response_url
143 .ok_or_else(|| ImageGenError::BadResponse("submit returned no response_url".into()))?;
144
145 let deadline = Instant::now() + MAX_WAIT;
147 loop {
148 let st: FalStatus = self.authed_get(&status_url).await?;
149 match st.status.as_str() {
150 "COMPLETED" => {
151 if let Some(err) = st.error {
153 return Err(ImageGenError::BadResponse(format!("fal job failed: {err}")));
154 }
155 break;
156 }
157 "FAILED" | "ERROR" => {
158 return Err(ImageGenError::BadResponse(format!(
159 "fal job failed: {}",
160 st.error.unwrap_or_else(|| "unknown".into())
161 )));
162 }
163 _ => {}
164 }
165 if Instant::now() >= deadline {
166 return Err(ImageGenError::Timeout {
167 context: "fal generation".into(),
168 });
169 }
170 tokio::time::sleep(POLL_INTERVAL).await;
171 }
172
173 let result: FalResult = self.authed_get(&response_url).await?;
175 let images = normalize_fal_images(&result.images, self.store.as_ref())?;
176
177 Ok(ImageGenResult {
178 images,
179 provider: "fal".into(),
180 model: model.into(),
181 revised_prompt: result.revised_prompt,
182 })
183 }
184}
185
186#[derive(Deserialize)]
189struct FalSubmit {
190 #[serde(default)]
191 #[allow(dead_code)]
192 request_id: Option<String>,
193 #[serde(default)]
194 status_url: Option<String>,
195 #[serde(default)]
196 response_url: Option<String>,
197}
198
199#[derive(Deserialize)]
200struct FalStatus {
201 status: String,
202 #[serde(default)]
203 error: Option<String>,
204 #[serde(default)]
205 #[allow(dead_code)]
206 error_type: Option<String>,
207}
208
209#[derive(Deserialize)]
210struct FalResult {
211 #[serde(default)]
212 images: Vec<FalImage>,
213 #[serde(default)]
214 revised_prompt: Option<String>,
215}
216
217#[derive(Deserialize)]
218struct FalImage {
219 url: Option<String>,
220 #[serde(default, rename = "b64_json")]
221 b64_json: Option<String>,
222 #[serde(default)]
223 width: Option<u32>,
224 #[serde(default)]
225 height: Option<u32>,
226}
227
228fn build_input(req: &ImageGenRequest) -> serde_json::Value {
232 let mut input = serde_json::json!({
233 "prompt": req.prompt,
234 "num_images": req.n.clamp(1, 8),
235 });
236 if let Some(size) = req.size {
237 let (w, h) = match size {
238 ImageSize::Square1024 => (1024, 1024),
239 ImageSize::Landscape1792 => (1792, 1024),
240 ImageSize::Portrait1792 => (1024, 1792),
241 };
242 input["image_size"] = serde_json::json!({ "width": w, "height": h });
243 }
244 if let Some(ref url) = req.reference_image_url {
245 input["image_url"] = serde_json::json!(url);
246 }
247 input
248}
249
250fn normalize_fal_images(
252 images: &[FalImage],
253 store: &dyn ImageSink,
254) -> Result<Vec<GeneratedImage>, ImageGenError> {
255 if images.is_empty() {
256 return Err(ImageGenError::BadResponse(
257 "fal result has no images".into(),
258 ));
259 }
260 images
261 .iter()
262 .map(|img| match (&img.url, &img.b64_json) {
263 (Some(url), _) => Ok(GeneratedImage {
264 url: url.clone(),
265 width: img.width,
266 height: img.height,
267 }),
268 (None, Some(b64)) => {
269 let bytes = base64::engine::general_purpose::STANDARD
270 .decode(b64)
271 .map_err(|e| ImageGenError::Base64(e.to_string()))?;
272 let url = store.save(bytes, "png")?;
273 Ok(GeneratedImage {
274 url,
275 width: img.width,
276 height: img.height,
277 })
278 }
279 (None, None) => Err(ImageGenError::BadResponse(
280 "fal image has neither url nor b64_json".into(),
281 )),
282 })
283 .collect()
284}
285
286#[cfg(test)]
287mod tests {
288 use parking_lot::Mutex;
289
290 use super::*;
291 use crate::image_gen::ImageSink;
292
293 struct FakeSink {
294 saved: Mutex<Vec<(Vec<u8>, String)>>,
295 }
296 impl ImageSink for FakeSink {
297 fn save(&self, bytes: Vec<u8>, ext: &str) -> Result<String, ImageGenError> {
298 let n = self.saved.lock().len();
299 let url = format!("/api/images/fal-{n}.{ext}");
300 self.saved.lock().push((bytes, ext.into()));
301 Ok(url)
302 }
303 }
304 fn sink() -> FakeSink {
305 FakeSink {
306 saved: Mutex::new(vec![]),
307 }
308 }
309 fn img(url: Option<&str>, b64: Option<&str>) -> FalImage {
310 FalImage {
311 url: url.map(Into::into),
312 b64_json: b64.map(Into::into),
313 width: Some(1024),
314 height: Some(1024),
315 }
316 }
317
318 #[test]
319 fn build_input_maps_prompt_num_size_and_reference() {
320 let req = ImageGenRequest {
321 prompt: "a cat".into(),
322 model: Some("fal-ai/flux/dev".into()),
323 n: 3,
324 size: Some(ImageSize::Landscape1792),
325 quality: None,
326 reference_image_url: Some("https://x/ref.png".into()),
327 };
328 let input = build_input(&req);
329 assert_eq!(input["prompt"], "a cat");
330 assert_eq!(input["num_images"], 3);
331 assert_eq!(input["image_size"]["width"], 1792);
332 assert_eq!(input["image_size"]["height"], 1024);
333 assert_eq!(input["image_url"], "https://x/ref.png");
334 }
335
336 #[test]
337 fn url_images_pass_through_with_dimensions() {
338 let s = sink();
339 let r = normalize_fal_images(&[img(Some("https://fal/cdn/1.png"), None)], &s).unwrap();
340 assert_eq!(r[0].url, "https://fal/cdn/1.png");
341 assert_eq!(r[0].width, Some(1024));
342 assert!(s.saved.lock().is_empty());
343 }
344
345 #[test]
346 fn b64_images_are_persisted() {
347 let s = sink();
348 let b64 = base64::engine::general_purpose::STANDARD.encode(b"falimg");
349 let r = normalize_fal_images(&[img(None, Some(&b64))], &s).unwrap();
350 assert!(r[0].url.starts_with("/api/images/fal-0.png"));
351 assert_eq!(s.saved.lock()[0].0, b"falimg");
352 }
353
354 #[test]
355 fn empty_images_is_an_error() {
356 let s = sink();
357 assert!(normalize_fal_images(&[], &s).is_err());
358 }
359
360 #[test]
361 fn submit_url_places_model_directly_in_path() {
362 let p = FalImageProvider {
363 client: reqwest::Client::new(),
364 base_url: "https://queue.fal.run".into(),
365 api_key: "k".into(),
366 store: Arc::new(crate::image_gen::FsImageStore::new(
367 std::path::PathBuf::from("/tmp"),
368 "/api/images/".into(),
369 )),
370 };
371 assert_eq!(
372 p.submit_url("fal-ai/flux/schnell"),
373 "https://queue.fal.run/fal-ai/flux/schnell"
374 );
375 assert_eq!(
376 p.submit_url("/fal-ai/flux/schnell"),
377 "https://queue.fal.run/fal-ai/flux/schnell"
378 );
379 }
380}