1use 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
17const DEFAULT_MODEL: &str = "black-forest-labs/flux-1-dev";
19
20const OPENROUTER_BASE_URL: &str = "https://openrouter.ai/api/v1";
22
23#[derive(Deserialize, JsonSchema)]
25pub struct GenerateImageArgs {
26 prompt: String,
27 model: Option<String>,
28 size: Option<String>,
29 n: Option<u64>,
30}
31
32const MAX_PROMPT_LEN: usize = 4000;
34
35pub struct GenerateImageTool;
37
38impl GenerateImageTool {
39 pub fn new() -> Self {
41 Self
42 }
43
44 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 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 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 let parsed: serde_json::Value =
112 serde_json::from_str(&text).map_err(|e| format!("Invalid JSON response: {}", e))?;
113
114 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 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 else if let Some(url_str) = item.get("url").and_then(|v| v.as_str()) {
132 images.push(url_str.as_bytes().to_vec());
134 }
135
136 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
158fn 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 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}