Skip to main content

oxi_agent/tools/
generate_image.rs

1//! Image generation tool using OpenRouter API.
2//!
3//! Provides an `AgentTool` that calls OpenRouter's image generation endpoint.
4//! Supports models like `black-forest-labs/flux-1-dev`, `openai/dall-e-3`, etc.
5
6use super::http_client::shared_http_client;
7use super::{AgentTool, AgentToolResult, ToolContext, ToolError};
8use crate::tools::typed::TypedTool;
9use async_trait::async_trait;
10use base64::{Engine, engine::general_purpose};
11use oxi_ai::types::{ImageGenerationRequest, ImageGenerationResponse};
12use schemars::JsonSchema;
13use serde::Deserialize;
14use serde_json::{Value, json};
15use std::env;
16
17/// Default image generation model.
18const DEFAULT_MODEL: &str = "black-forest-labs/flux-1-dev";
19
20/// OpenRouter API base URL.
21const OPENROUTER_BASE_URL: &str = "https://openrouter.ai/api/v1";
22
23/// Arguments for image generation.
24#[derive(Deserialize, JsonSchema)]
25pub struct GenerateImageArgs {
26    prompt: String,
27    model: Option<String>,
28    size: Option<String>,
29    n: Option<u64>,
30}
31
32/// Maximum prompt length to warn about.
33const MAX_PROMPT_LEN: usize = 4000;
34
35/// Image generation tool.
36pub struct GenerateImageTool;
37
38impl GenerateImageTool {
39    /// Create a new GenerateImageTool.
40    pub fn new() -> Self {
41        Self
42    }
43
44    /// Build the request body for OpenRouter.
45    fn build_request_body(req: &ImageGenerationRequest) -> serde_json::Value {
46        let mut body = serde_json::json!({
47            "model": req.model.as_deref().unwrap_or(DEFAULT_MODEL),
48            "prompt": req.prompt,
49        });
50
51        if let Some(size) = &req.size {
52            body["size"] = serde_json::json!(size);
53        }
54        if let Some(n) = req.n {
55            body["n"] = serde_json::json!(n);
56        }
57        if let Some(ref fmt) = req.response_format {
58            body["response_format"] = serde_json::json!(fmt);
59        }
60        body
61    }
62
63    /// Call OpenRouter image generation API.
64    async fn call_openrouter(
65        &self,
66        api_key: &str,
67        request: &ImageGenerationRequest,
68    ) -> Result<ImageGenerationResponse, ToolError> {
69        let url = format!("{}/images/generations", OPENROUTER_BASE_URL);
70        let body = Self::build_request_body(request);
71
72        let client = shared_http_client();
73        let resp = client
74            .post(&url)
75            .header("Authorization", format!("Bearer {}", api_key))
76            .header("Content-Type", "application/json")
77            .header("HTTP-Referer", "https://github.com/oxi")
78            .json(&body)
79            .send()
80            .await
81            .map_err(|e| format!("OpenRouter request failed: {}", e))?;
82
83        let status = resp.status();
84        let text = resp
85            .text()
86            .await
87            .map_err(|e| format!("Failed to read response: {}", e))?;
88
89        if !status.is_success() {
90            // Try to extract error message from API response
91            let err_msg = {
92                let parsed = serde_json::from_str::<serde_json::Value>(&text).ok();
93                match parsed {
94                    Some(ref root) => root
95                        .get("error")
96                        .or_else(|| root.get("message"))
97                        .and_then(|v| v.as_str())
98                        .map(String::from)
99                        .unwrap_or_else(|| text.clone()),
100                    None => text.clone(),
101                }
102            };
103            return Err(format!(
104                "OpenRouter API error ({}): {}",
105                status,
106                err_msg.clone()
107            ));
108        }
109
110        // Parse OpenRouter image response
111        let parsed: serde_json::Value =
112            serde_json::from_str(&text).map_err(|e| format!("Invalid JSON response: {}", e))?;
113
114        // OpenRouter wraps the standard response under "data"
115        let data = parsed
116            .get("data")
117            .ok_or_else(|| "Missing 'data' field in response".to_string())?
118            .as_array()
119            .ok_or_else(|| "'data' is not an array".to_string())?;
120
121        let mut images: Vec<Vec<u8>> = Vec::new();
122        let mut revised_prompt: Option<String> = None;
123
124        for item in data {
125            // b64_json format
126            if let Some(b64) = item.get("b64_json").and_then(|v| v.as_str()) {
127                let bytes = base64_decode(b64)?;
128                images.push(bytes);
129            }
130            // url format (silently included as-is; return URL as string)
131            else if let Some(url_str) = item.get("url").and_then(|v| v.as_str()) {
132                // Encode URL as bytes for uniform output
133                images.push(url_str.as_bytes().to_vec());
134            }
135
136            // Capture revised_prompt if present (DALL-E style)
137            if revised_prompt.is_none() {
138                revised_prompt = item
139                    .get("revised_prompt")
140                    .and_then(|v| v.as_str())
141                    .map(String::from);
142            }
143        }
144
145        Ok(ImageGenerationResponse {
146            images,
147            revised_prompt,
148        })
149    }
150}
151
152impl Default for GenerateImageTool {
153    fn default() -> Self {
154        Self::new()
155    }
156}
157
158/// Decode base64 without the full base64 crate — plain std.
159fn base64_decode(input: &str) -> Result<Vec<u8>, ToolError> {
160    const CHARS: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
161    let input = input.as_bytes();
162    let mut out = Vec::with_capacity(input.len() * 3 / 4);
163    let mut buf: u32 = 0;
164    let mut bits = 0;
165
166    for &byte in input
167        .iter()
168        .filter(|&&b| b != b'=' && b != b'\n' && b != b'\r')
169    {
170        let val = CHARS
171            .iter()
172            .position(|&c| c == byte)
173            .ok_or_else(|| format!("Invalid base64 character: {:?}", byte as char))?
174            as u32;
175        buf = (buf << 6) | val;
176        bits += 6;
177        if bits >= 8 {
178            bits -= 8;
179            out.push((buf >> bits) as u8);
180            buf &= (1 << bits) - 1;
181        }
182    }
183    Ok(out)
184}
185
186#[async_trait]
187impl AgentTool for GenerateImageTool {
188    fn name(&self) -> &str {
189        "generate_image"
190    }
191
192    fn label(&self) -> &str {
193        "Generate Image"
194    }
195
196    fn description(&self) -> &str {
197        "Generate an image from a text prompt using an AI image generation model via OpenRouter. \
198         Takes a `prompt` (required), optional `model` (default: black-forest-labs/flux-1-dev), \
199         and optional `size`. Returns base64-encoded image data."
200    }
201
202    fn parameters_schema(&self) -> Value {
203        json!({
204            "type": "object",
205            "properties": {
206                "prompt": {
207                    "type": "string",
208                    "description": "Detailed text description of the desired image"
209                },
210                "model": {
211                    "type": "string",
212                    "description": "Model to use (e.g. openai/dall-e-3, black-forest-labs/flux-1-dev, stability-ai/stable-diffusion-3). Default: black-forest-labs/flux-1-dev"
213                },
214                "size": {
215                    "type": "string",
216                    "description": "Image size (e.g. 1024x1024, 1024x1792). Provider-dependent."
217                },
218                "n": {
219                    "type": "integer",
220                    "minimum": 1,
221                    "maximum": 10,
222                    "description": "Number of images to generate (1-10, default 1)"
223                }
224            },
225            "required": ["prompt"]
226        })
227    }
228
229    async fn execute(
230        &self,
231        _tool_call_id: &str,
232        params: Value,
233        _signal: Option<tokio::sync::oneshot::Receiver<()>>,
234        _ctx: &ToolContext,
235    ) -> Result<AgentToolResult, ToolError> {
236        let args: GenerateImageArgs =
237            serde_json::from_value(params).map_err(|e| format!("invalid params: {e}"))?;
238        self.execute_typed(_tool_call_id, args, _signal, _ctx).await
239    }
240}
241
242#[async_trait]
243impl TypedTool for GenerateImageTool {
244    type Args = GenerateImageArgs;
245
246    async fn execute_typed(
247        &self,
248        _tool_call_id: &str,
249        args: Self::Args,
250        _signal: Option<tokio::sync::oneshot::Receiver<()>>,
251        _ctx: &ToolContext,
252    ) -> Result<AgentToolResult, ToolError> {
253        if args.prompt.is_empty() {
254            return Err("Prompt cannot be empty".to_string());
255        }
256        if args.prompt.chars().count() > MAX_PROMPT_LEN {
257            tracing::warn!(
258                "Prompt length {} exceeds recommended max {}",
259                args.prompt.chars().count(),
260                MAX_PROMPT_LEN
261            );
262        }
263        let request = ImageGenerationRequest {
264            prompt: args.prompt,
265            model: args.model,
266            size: args.size,
267            n: args.n.map(|v| v as u32),
268            response_format: Some("b64_json".to_string()),
269        };
270        let api_key = env::var("OPENROUTER_API_KEY")
271            .or_else(|_| env::var("OPENAI_API_KEY"))
272            .map_err(|_| "OPENROUTER_API_KEY (or OPENAI_API_KEY) environment variable is not set. Please set your API key before using the image generation tool.")?;
273        let response = self.call_openrouter(&api_key, &request).await?;
274        if response.images.is_empty() {
275            return Ok(AgentToolResult::success(
276                "Image generation completed but returned no images.",
277            ));
278        }
279        let n_images = response.images.len();
280        let mut output = format!("Generated {} image(s).\n\n", n_images);
281        if let Some(ref revised) = response.revised_prompt {
282            output.push_str(&format!("Revised prompt: {}\n\n", revised));
283        }
284        for (i, img_data) in response.images.iter().enumerate() {
285            let b64 = general_purpose::STANDARD.encode(img_data);
286            output.push_str(&format!(
287                "Image {} ({} bytes, base64):\n{}\n\n",
288                i + 1,
289                img_data.len(),
290                b64
291            ));
292        }
293        Ok(AgentToolResult::success(output.trim_end()))
294    }
295}
296
297#[cfg(test)]
298mod tests {
299    use super::*;
300
301    #[test]
302    fn test_base64_decode() {
303        // "Hello, World!" base64-encoded
304        let encoded = "SGVsbG8sIFdvcmxkIQ==";
305        let decoded = base64_decode(encoded).unwrap();
306        assert_eq!(decoded, b"Hello, World!");
307    }
308
309    #[test]
310    fn test_build_request_body() {
311        let req = ImageGenerationRequest {
312            prompt: "A red cat".to_string(),
313            model: Some("flux-dev".to_string()),
314            size: Some("1024x1024".to_string()),
315            n: Some(2),
316            response_format: Some("b64_json".to_string()),
317        };
318
319        let body = GenerateImageTool::build_request_body(&req);
320        assert_eq!(body["prompt"], "A red cat");
321        assert_eq!(body["model"], "flux-dev");
322        assert_eq!(body["size"], "1024x1024");
323        assert_eq!(body["n"], 2);
324        assert_eq!(body["response_format"], "b64_json");
325    }
326
327    #[test]
328    fn test_default_model() {
329        let req = ImageGenerationRequest::default();
330        let body = GenerateImageTool::build_request_body(&req);
331        assert_eq!(body["model"], DEFAULT_MODEL);
332    }
333
334    #[test]
335    fn test_image_generation_response_default() {
336        let resp = ImageGenerationResponse::default();
337        assert!(resp.images.is_empty());
338        assert!(resp.revised_prompt.is_none());
339    }
340}