Skip to main content

oxios_kernel/image_gen/
fal.rs

1//! fal.ai image-generation provider (queue API: submit → poll → fetch).
2//!
3//! fal models are asynchronous: submit returns a `request_id` plus absolute
4//! `status_url`/`response_url`; poll `status_url` until `COMPLETED`, then GET
5//! `response_url`. This mirrors LobeHub's `fal.subscribe()` (which blocks
6//! until completion). Verified against the fal queue REST docs
7//! (<https://fal.ai/docs/documentation/model-apis/inference/queue>):
8//!
9//! - Host: `https://queue.fal.run`
10//! - Submit: `POST {base}/{model_id}` — body is the input object **directly**
11//!   (NOT wrapped in `{"input": ...}`; the JS SDK wraps client-side, the REST
12//!   API does not).
13//! - Submit response: `{request_id, status_url, response_url, ...}` (absolute
14//!   URLs) — use them verbatim, sidestepping model_id-in-path construction.
15//! - Status: `GET {status_url}` → `{status, error?, error_type?}`. `COMPLETED`
16//!   WITH an `error` field means the request FAILED.
17//! - Result: `GET {response_url}` → model-specific, images at TOP LEVEL
18//!   (`{images: [{url, width, height, content_type}], ...}`).
19//! - Auth: `Authorization: Key {FAL_KEY}`.
20
21use 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
33/// Default fal queue host.
34pub const FAL_DEFAULT_BASE: &str = "https://queue.fal.run";
35/// Poll interval for fal queue status (matches LobeHub `WAIT_POLL_INTERVAL_MS`).
36const POLL_INTERVAL: Duration = Duration::from_millis(3000);
37/// Max wait for a fal generation (matches LobeHub `MAX_WAIT_TIMEOUT_MS`).
38const MAX_WAIT: Duration = Duration::from_millis(175_000);
39
40/// fal.ai provider via the queue REST API.
41pub struct FalImageProvider {
42    client: reqwest::Client,
43    base_url: String,
44    api_key: String,
45    store: Arc<dyn ImageSink>,
46}
47
48impl FalImageProvider {
49    /// Construct from already-resolved values (the tool layer resolves the
50    /// `fal` credential and image dir).
51    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    /// Submit URL: `{base}/{model}` (model carries its own path, e.g.
74    /// `fal-ai/flux/schnell`).
75    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        // 1. Submit — body is the input object DIRECTLY (not wrapped).
118        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        // 2. Poll status until COMPLETED / FAILED / timeout.
146        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                    // COMPLETED may still carry an error → the job failed.
152                    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        // 3. Fetch result — images at TOP LEVEL (no `.data` wrapper).
174        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// ── Wire types ────────────────────────────────────────────────────────────
187
188#[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
228// ── Pure helpers (unit-tested) ────────────────────────────────────────────
229
230/// Build the fal input payload (sent DIRECTLY as the submit body).
231fn 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
250/// Normalize fal output images to fetchable URLs (carrying dimensions).
251fn 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}