oxios_kernel/tools/builtin/
image_generation_tool.rs1use 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
29const IMAGE_SERVE_PREFIX: &str = "/api/images/";
31
32pub struct ImageGenerationTool {
34 provider: String,
35 base_url: String,
36 default_model: String,
37 default_num: u8,
38 engine_api_key: Option<String>,
40 images_dir: PathBuf,
41}
42
43impl ImageGenerationTool {
44 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 `` 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 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 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
259fn 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
269fn 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#[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}