Skip to main content

oxios_kernel/tools/builtin/
image_generation_tool.rs

1//! Image generation tool — wraps the `image_gen` provider behind `AgentTool`.
2//!
3//! Action-based (mirrors LobeHub's 4-API builtin tool, pared down to the
4//! synchronous Phase-1 surface):
5//! - `generate` — generate images from a prompt.
6//! - `list_models` — echo the configured provider/model (availability check).
7//!
8//! The provider, base URL, and default model come from `[image-gen]` config.
9//! The API key is resolved via [`CredentialStore`] using the same chain as
10//! the chat engine — no separate credential is needed.
11//!
12//! Config is read once at construction (a boot-time snapshot via
13//! `kernel.infra.config()`); the tool registers per-agent, so Phase-1 does
14//! not need hot-reload.
15
16use std::path::PathBuf;
17
18use async_trait::async_trait;
19use oxi_sdk::{AgentTool, AgentToolResult, ToolContext};
20use serde_json::{Value, json};
21
22use crate::credential::CredentialStore;
23use crate::image_gen::{
24    FalImageProvider, FsImageStore, GeneratedImage, ImageGenProvider, ImageGenRequest, ImageSize,
25    OpenAiImageProvider,
26};
27use crate::kernel_handle::KernelHandle;
28
29/// URL prefix under which persisted images are served (see `image_routes.rs`).
30const IMAGE_SERVE_PREFIX: &str = "/api/images/";
31
32/// Agent tool for image generation (Phase 1: OpenAI-compatible providers).
33pub struct ImageGenerationTool {
34    provider: String,
35    base_url: String,
36    default_model: String,
37    default_num: u8,
38    /// `[engine].api_key` override forwarded to credential resolution.
39    engine_api_key: Option<String>,
40    images_dir: PathBuf,
41}
42
43impl ImageGenerationTool {
44    /// Create from a [`KernelHandle`].
45    ///
46    /// Reads the `[image-gen]` config snapshot once and resolves the
47    /// workspace `images/` dir for persisted results.
48    pub fn from_kernel(kernel: &KernelHandle) -> Self {
49        let cfg = kernel.infra.config();
50        let ig = &cfg.image_gen;
51        Self {
52            provider: ig.provider.clone(),
53            base_url: ig.base_url.clone(),
54            default_model: ig.default_model.clone(),
55            default_num: ig.default_num,
56            engine_api_key: cfg.api_key(),
57            images_dir: kernel.state.workspace_path().join("images"),
58        }
59    }
60}
61
62impl std::fmt::Debug for ImageGenerationTool {
63    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
64        f.debug_struct("ImageGenerationTool")
65            .field("provider", &self.provider)
66            .field("images_dir", &self.images_dir)
67            .finish()
68    }
69}
70
71#[async_trait]
72impl AgentTool for ImageGenerationTool {
73    fn name(&self) -> &str {
74        "image_generation"
75    }
76
77    fn label(&self) -> &str {
78        "Image Generation"
79    }
80
81    fn description(&self) -> &'static str {
82        "Generate images from a text prompt via an OpenAI-compatible image model. \
83         When generation completes, show each image by emitting markdown \
84         `![](url)` using the URLs from the result EXACTLY as given — do not \
85         rewrite, shorten, or translate them. Include a brief caption only. \
86         Do not retry automatically on content-policy or billing errors; \
87         report the error concisely instead."
88    }
89
90    fn parameters_schema(&self) -> Value {
91        json!({
92            "type": "object",
93            "properties": {
94                "action": {
95                    "type": "string",
96                    "enum": ["generate", "list_models"],
97                    "description": "Operation: 'generate' (default) or 'list_models' (show configured model)."
98                },
99                "prompt": {
100                    "type": "string",
101                    "description": "Text-to-image prompt (required for 'generate')."
102                },
103                "model": {
104                    "type": "string",
105                    "description": "Provider model id. Omit to use the configured default."
106                },
107                "n": {
108                    "type": "integer",
109                    "minimum": 1,
110                    "maximum": 8,
111                    "default": 1,
112                    "description": "Number of images to generate."
113                },
114                "size": {
115                    "type": "string",
116                    "enum": ["1024x1024", "1792x1024", "1024x1792"],
117                    "description": "Output dimensions. Omit for provider default."
118                },
119                "quality": {
120                    "type": "string",
121                    "description": "Quality hint (e.g. 'standard', 'hd'). Provider-specific."
122                },
123                "reference_image_url": {
124                    "type": "string",
125                    "description": "Reference image URL for image-to-image (fal providers). Optional."
126                }
127            },
128            "required": ["action"]
129        })
130    }
131
132    async fn execute(
133        &self,
134        _tool_call_id: &str,
135        params: Value,
136        _signal: Option<tokio::sync::oneshot::Receiver<()>>,
137        _ctx: &ToolContext,
138    ) -> Result<AgentToolResult, oxi_sdk::ToolError> {
139        let action = params
140            .get("action")
141            .and_then(|v| v.as_str())
142            .unwrap_or("generate");
143
144        if action == "list_models" {
145            return Ok(AgentToolResult::success(
146                serde_json::to_string_pretty(&json!({
147                    "provider": self.provider,
148                    "default_model": self.default_model,
149                    "default_num": self.default_num,
150                    "sizes": ["1024x1024", "1792x1024", "1024x1792"],
151                }))
152                .unwrap_or_default(),
153            ));
154        }
155
156        if action != "generate" {
157            return Err(format!(
158                "Unknown action '{action}'. Valid: generate, list_models."
159            ));
160        }
161
162        let prompt = params
163            .get("prompt")
164            .and_then(|v| v.as_str())
165            .ok_or_else(|| "Missing required parameter: prompt".to_string())?;
166
167        // Model: explicit param → configured default → error.
168        let model = params
169            .get("model")
170            .and_then(|v| v.as_str())
171            .map(str::to_owned)
172            .or_else(|| {
173                if self.default_model.is_empty() {
174                    None
175                } else {
176                    Some(self.default_model.clone())
177                }
178            })
179            .ok_or_else(|| {
180                "No model specified and no [image-gen].default_model configured".to_string()
181            })?;
182
183        let n = params
184            .get("n")
185            .and_then(|v| v.as_u64())
186            .unwrap_or(u64::from(self.default_num)) as u8;
187
188        let size = params
189            .get("size")
190            .and_then(|v| v.as_str())
191            .and_then(parse_size);
192
193        let quality = params
194            .get("quality")
195            .and_then(|v| v.as_str())
196            .map(str::to_owned);
197
198        // Resolve the API key via the same chain the chat engine uses.
199        let api_key = match CredentialStore::resolve(&self.provider, self.engine_api_key.as_deref())
200        {
201            Some((key, _src)) => key,
202            None => {
203                return Ok(AgentToolResult::error(format!(
204                    "No API key resolved for provider '{}'. Set it via the engine key, \
205                         ~/.oxios/auth.json, or OXIOS_{}_API_KEY.",
206                    self.provider,
207                    self.provider.to_uppercase()
208                )));
209            }
210        };
211
212        let store = std::sync::Arc::new(FsImageStore::new(
213            self.images_dir.clone(),
214            IMAGE_SERVE_PREFIX.into(),
215        )) as std::sync::Arc<dyn crate::image_gen::ImageSink>;
216
217        let provider: Box<dyn ImageGenProvider> = match self.provider.as_str() {
218            "fal" => {
219                let fal_base = fal_base_url(&self.base_url);
220                match FalImageProvider::new(fal_base, api_key.clone(), store) {
221                    Ok(p) => Box::new(p),
222                    Err(e) => {
223                        return Ok(AgentToolResult::error(format!("provider init failed: {e}")));
224                    }
225                }
226            }
227            _ => match OpenAiImageProvider::new(self.base_url.clone(), api_key, store) {
228                Ok(p) => Box::new(p),
229                Err(e) => return Ok(AgentToolResult::error(format!("provider init failed: {e}"))),
230            },
231        };
232
233        let reference_image_url = params
234            .get("reference_image_url")
235            .and_then(|v| v.as_str())
236            .map(str::to_owned);
237
238        let req = ImageGenRequest {
239            prompt: prompt.to_owned(),
240            model: Some(model),
241            n,
242            size,
243            quality,
244            reference_image_url,
245        };
246
247        match provider.generate(&req).await {
248            Ok(result) => Ok(AgentToolResult::success(
249                serde_json::to_string(&GenerationToolOutput::new(result, prompt))
250                    .unwrap_or_default(),
251            )),
252            Err(e) => Ok(AgentToolResult::error(format!(
253                "image generation failed: {e}"
254            ))),
255        }
256    }
257}
258
259/// Parse an OpenAI size string into [`ImageSize`].
260fn parse_size(s: &str) -> Option<ImageSize> {
261    match s {
262        "1024x1024" => Some(ImageSize::Square1024),
263        "1792x1024" => Some(ImageSize::Landscape1792),
264        "1024x1792" => Some(ImageSize::Portrait1792),
265        _ => None,
266    }
267}
268
269/// Pick the fal queue base URL: use the configured base unless it is empty or
270/// still the OpenAI default (common when the user switches provider to fal).
271fn fal_base_url(configured: &str) -> String {
272    let c = configured.trim();
273    if c.is_empty() || c.contains("openai.com") {
274        crate::image_gen::FAL_DEFAULT_BASE.into()
275    } else {
276        c.into()
277    }
278}
279
280/// Tool output payload (serialized as the tool result string).
281#[derive(serde::Serialize)]
282struct GenerationToolOutput {
283    action: &'static str,
284    images: Vec<GeneratedImage>,
285    prompt: String,
286    provider: String,
287    model: String,
288    revised_prompt: Option<String>,
289}
290
291impl GenerationToolOutput {
292    fn new(r: crate::image_gen::ImageGenResult, prompt: &str) -> Self {
293        Self {
294            action: "generate",
295            prompt: prompt.to_owned(),
296            images: r.images,
297            provider: r.provider,
298            model: r.model,
299            revised_prompt: r.revised_prompt,
300        }
301    }
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307
308    fn tool() -> ImageGenerationTool {
309        ImageGenerationTool {
310            provider: "openai".into(),
311            base_url: "https://api.openai.com/v1".into(),
312            default_model: "gpt-image-1".into(),
313            default_num: 1,
314            engine_api_key: None,
315            images_dir: PathBuf::from("/tmp"),
316        }
317    }
318
319    #[test]
320    fn parse_size_maps_known_strings() {
321        assert_eq!(parse_size("1024x1024"), Some(ImageSize::Square1024));
322        assert_eq!(parse_size("1792x1024"), Some(ImageSize::Landscape1792));
323        assert_eq!(parse_size("1024x1792"), Some(ImageSize::Portrait1792));
324        assert_eq!(parse_size("bogus"), None);
325    }
326
327    #[test]
328    fn schema_has_generate_and_list_models() {
329        let schema = tool().parameters_schema();
330        let actions = schema["properties"]["action"]["enum"].as_array().unwrap();
331        assert!(actions.iter().any(|a| a == "generate"));
332        assert!(actions.iter().any(|a| a == "list_models"));
333    }
334}