Skip to main content

robit_agent/tool/
generate_image.rs

1//! `generate_image` tool - generates images from text prompts.
2//!
3//! Uses the configured `default_image_model` provider (Wanxiang/DashScope or
4//! any OpenAI-compatible image API). The model is configured server-side and
5//! is not exposed to the LLM. Generated images are downloaded and saved to
6//! disk; the tool returns a JSON summary with saved paths and source URLs.
7
8use async_trait::async_trait;
9use serde::Deserialize;
10use serde_json::{json, Value};
11use std::path::Path;
12use time::macros::format_description;
13use time::OffsetDateTime;
14
15use super::async_runner::AsyncTaskWork;
16use super::{resolve_path, Tool, ToolContext, ToolResult};
17use crate::error::Result;
18use crate::image_gen::{ImageGenClient, ImageGenRequest};
19use crate::media::download_media;
20
21/// Maximum number of images that can be generated in one call.
22const MAX_N: u32 = 4;
23
24/// Maximum seed value accepted by the DashScope (Wanxiang) API.
25const MAX_SEED: u64 = 2_147_483_647;
26
27#[derive(Debug, Deserialize)]
28struct GenerateImageArgs {
29    prompt: String,
30    #[serde(default)]
31    filename: Option<String>,
32    #[serde(default)]
33    output_path: Option<String>,
34    #[serde(default)]
35    n: Option<u32>,
36    /// Output resolution as "width*height" (e.g. "1280*1280").
37    #[serde(default)]
38    size: Option<String>,
39    /// Negative prompt: content to exclude from the image.
40    #[serde(default)]
41    negative_prompt: Option<String>,
42    /// Whether to enable smart prompt rewriting (provider default: true).
43    #[serde(default)]
44    prompt_extend: Option<bool>,
45    /// Random seed in [0, 2147483647] for reproducible generation.
46    #[serde(default)]
47    seed: Option<u64>,
48}
49
50/// Build the provider pass-through parameters from the optional tool args.
51///
52/// Returns `Value::Null` when no optional parameter was given, so the request
53/// body stays identical to before these options existed. For the DashScope
54/// protocol these keys are merged into `parameters`; for the OpenAI protocol
55/// into the top-level request body (unsupported keys are rejected by the
56/// provider, which surfaces as an API error the LLM can react to).
57fn build_extra_params(args: &GenerateImageArgs) -> Value {
58    let mut extra = serde_json::Map::new();
59    if let Some(size) = args
60        .size
61        .as_deref()
62        .map(str::trim)
63        .filter(|s| !s.is_empty())
64    {
65        extra.insert("size".to_string(), json!(size));
66    }
67    if let Some(np) = args
68        .negative_prompt
69        .as_deref()
70        .map(str::trim)
71        .filter(|s| !s.is_empty())
72    {
73        extra.insert("negative_prompt".to_string(), json!(np));
74    }
75    if let Some(pe) = args.prompt_extend {
76        extra.insert("prompt_extend".to_string(), json!(pe));
77    }
78    if let Some(seed) = args.seed {
79        extra.insert("seed".to_string(), json!(seed));
80    }
81    if extra.is_empty() {
82        Value::Null
83    } else {
84        Value::Object(extra)
85    }
86}
87
88pub struct GenerateImageTool {
89    client: ImageGenClient,
90}
91
92impl GenerateImageTool {
93    pub fn new(client: ImageGenClient) -> Self {
94        Self { client }
95    }
96}
97
98#[async_trait]
99impl Tool for GenerateImageTool {
100    fn name(&self) -> &str {
101        "generate_image"
102    }
103
104    fn description(&self) -> &str {
105        "Generate images from a text prompt using AI image generation. \
106         The model is configured server-side and cannot be changed by the caller. \
107         Generated images are saved as PNG files and the paths are returned."
108    }
109
110    fn parameters_schema(&self) -> Value {
111        json!({
112            "type": "object",
113            "properties": {
114                "prompt": {
115                    "type": "string",
116                    "description": "Text description of the image to generate. Supports Chinese and English."
117                },
118                "filename": {
119                    "type": "string",
120                    "description": "Base filename (without extension) for saved images. \
121                                    If omitted, a timestamp-based name is generated. \
122                                    For multiple images, a '-1', '-2' suffix is appended."
123                },
124                "output_path": {
125                    "type": "string",
126                    "description": "Directory to save images (relative or absolute). \
127                                    Defaults to {working_dir}/images."
128                },
129                "n": {
130                    "type": "integer",
131                    "description": "Number of images to generate (1-4). Defaults to 1.",
132                    "minimum": 1,
133                    "maximum": MAX_N
134                },
135                "size": {
136                    "type": "string",
137                    "description": "Output image resolution as 'width*height' (e.g. '1280*1280'). \
138                                    Omit to use the provider default. Common ratios (Wanxiang wan2.5+): \
139                                    1:1 '1280*1280', 3:4 '1104*1472', 4:3 '1472*1104', \
140                                    9:16 '960*1696', 16:9 '1696*960'. \
141                                    Constraints depend on the configured model; an invalid size is \
142                                    rejected by the provider as an API error."
143                },
144                "negative_prompt": {
145                    "type": "string",
146                    "description": "Optional negative prompt: content to avoid in the generated \
147                                    image (e.g. '低分辨率,肢体畸形'). Max 500 characters."
148                },
149                "prompt_extend": {
150                    "type": "boolean",
151                    "description": "Optional. Enable smart prompt rewriting (provider default: true). \
152                                    Set to false if generation fails with IPInfringementSuspect or \
153                                    DataInspectionFailed caused by the rewritten prompt."
154                },
155                "seed": {
156                    "type": "integer",
157                    "description": "Optional random seed in [0, 2147483647]. Same seed keeps \
158                                    results relatively stable across calls.",
159                    "minimum": 0,
160                    "maximum": MAX_SEED
161                }
162            },
163            "required": ["prompt"]
164        })
165    }
166
167    fn requires_confirmation(&self) -> bool {
168        true
169    }
170
171    fn supports_async(&self) -> bool {
172        true
173    }
174
175    async fn execute(&self, args: Value, ctx: &ToolContext) -> Result<ToolResult> {
176        let parsed: GenerateImageArgs = match serde_json::from_value(args) {
177            Ok(a) => a,
178            Err(e) => return Ok(ToolResult::error(format!("Argument parsing failed: {}", e))),
179        };
180
181        if parsed.prompt.trim().is_empty() {
182            return Ok(ToolResult::error("prompt cannot be empty".to_string()));
183        }
184
185        // Validate and clamp n
186        let n = parsed.n.unwrap_or(1).clamp(1, MAX_N);
187
188        // Validate seed against the provider's accepted range (schema also
189        // declares it, but the LLM may still send an out-of-range value).
190        if let Some(seed) = parsed.seed {
191            if seed > MAX_SEED {
192                return Ok(ToolResult::error(format!(
193                    "seed must be in [0, {}], got {}",
194                    MAX_SEED, seed
195                )));
196            }
197        }
198
199        let extra_params = build_extra_params(&parsed);
200
201        // Resolve save directory (default: {working_dir}/images)
202        let save_dir = match parsed.output_path.as_deref() {
203            Some(p) => resolve_path(p, &ctx.working_dir),
204            None => ctx.working_dir.join("images"),
205        };
206
207        // Determine base filename (default: image_{YYYYMMDD_HHMMSS})
208        let base_filename = parsed
209            .filename
210            .as_deref()
211            .filter(|s| !s.trim().is_empty())
212            .map(|s| s.to_string())
213            .unwrap_or_else(default_filename);
214
215        // The actual generation + download can take 30-60s (or minutes for
216        // video), so it runs in a background task. We validate args above
217        // (cheap, gives immediate feedback on bad input) and move the heavy
218        // work into `work`, returning a pending placeholder.
219        let client = self.client.clone();
220        let working_dir = ctx.working_dir.clone();
221        let prompt = parsed.prompt.clone();
222
223        let work: AsyncTaskWork = Box::pin(async move {
224            let req = ImageGenRequest {
225                prompt,
226                n: Some(n),
227                extra_params,
228            };
229
230            tracing::info!(
231                "[generate_image] requesting {} image(s) (background), extra_params={}",
232                n,
233                req.extra_params
234            );
235
236            let images = match client.generate(&req).await {
237                Ok(imgs) => imgs,
238                Err(e) => {
239                    tracing::error!(
240                        "[generate_image] image generation failed: {}. \
241                         The error will be reported to the Agent as a task result.",
242                        e
243                    );
244                    let info = e.to_error_info();
245                    let err_json = json!({
246                        "status": "failed",
247                        "error": {
248                            "kind": info.kind,
249                            "code": info.code,
250                            "message": info.message,
251                            "retryable": info.retryable,
252                        }
253                    });
254                    return ToolResult::error(
255                        serde_json::to_string_pretty(&err_json)
256                            .unwrap_or_else(|_| err_json.to_string()),
257                    );
258                }
259            };
260
261            if images.is_empty() {
262                let err_json = json!({
263                    "status": "failed",
264                    "error": "Provider returned no images"
265                });
266                return ToolResult::error(
267                    serde_json::to_string_pretty(&err_json)
268                        .unwrap_or_else(|_| err_json.to_string()),
269                );
270            }
271
272            // Download and save each image. All images are attempted even if
273            // some fail, so partial results are preserved.
274            let multi = images.len() > 1;
275            let mut results: Vec<Value> = Vec::with_capacity(images.len());
276            let mut success_count: usize = 0;
277
278            for (i, img) in images.iter().enumerate() {
279                let index = i + 1;
280                let filename = if multi {
281                    format!("{}-{}.png", base_filename, index)
282                } else {
283                    format!("{}.png", base_filename)
284                };
285
286                let saved_path = download_media(&img.url, Some(&filename), &save_dir).await;
287                match saved_path {
288                    Ok(path) => {
289                        success_count += 1;
290                        results.push(json!({
291                            "index": index,
292                            "file": display_path(&path, &working_dir),
293                            "size": img.size.clone().unwrap_or_else(|| "unknown".to_string()),
294                            "url": img.url,
295                        }));
296                    }
297                    Err(e) => {
298                        results.push(json!({
299                            "index": index,
300                            "file": null,
301                            "size": img.size.clone().unwrap_or_else(|| "unknown".to_string()),
302                            "url": img.url,
303                            "error": format!("Download failed: {}", e),
304                        }));
305                    }
306                }
307            }
308
309            let status = if success_count == images.len() {
310                "success"
311            } else {
312                "partial"
313            };
314
315            let response = json!({
316                "status": status,
317                "generated_count": success_count,
318                "images": results,
319            });
320
321            let content = serde_json::to_string_pretty(&response)
322                .unwrap_or_else(|_| response.to_string());
323
324            if success_count == 0 {
325                // All downloads failed - report as error
326                ToolResult::error(content)
327            } else {
328                ToolResult::success(content)
329            }
330        });
331
332        // Submit the background task and return a placeholder. The Agent tracks
333        // the task id and reinjects the final result when `work` completes.
334        let task_id = ctx.async_runner.submit(
335            ctx.tool_call_id.clone(),
336            ctx.session_id.clone(),
337            self.name().to_string(),
338            work,
339            ctx.cancel_token.clone(),
340        );
341
342        let placeholder = format!(
343            "图片生成中(异步任务 task_id={})。预计耗时 30-60 秒,完成后会自动通知结果。\
344             你可以继续其他工作,完成后我会收到通知并告知你。",
345            task_id
346        );
347        Ok(ToolResult::pending(placeholder, task_id))
348    }
349}
350
351/// Generate a timestamp-based default filename: `image_{YYYYMMDD_HHMMSS}`.
352fn default_filename() -> String {
353    const FMT: &[time::format_description::FormatItem<'_>] =
354        format_description!("image_[year][month][day]_[hour][minute][second]");
355    OffsetDateTime::now_utc()
356        .format(FMT)
357        .unwrap_or_else(|_| "image".to_string())
358}
359
360/// Render a saved path relative to the working directory when possible,
361/// otherwise fall back to the absolute path.
362fn display_path(path: &Path, working_dir: &Path) -> String {
363    if let Ok(rel) = path.strip_prefix(working_dir) {
364        // Use forward slashes for display consistency across platforms.
365        rel.to_string_lossy().replace('\\', "/")
366    } else {
367        path.to_string_lossy().replace('\\', "/")
368    }
369}
370
371#[cfg(test)]
372mod tests {
373    use super::*;
374    use std::path::PathBuf;
375
376    #[test]
377    fn test_default_filename_format() {
378        let name = default_filename();
379        assert!(name.starts_with("image_"), "filename was: {name}");
380        // image_ + 8 digits + _ + 6 digits
381        assert!(name.len() >= "image_YYYYMMDD_HHMMSS".len(), "filename was: {name}");
382    }
383
384    #[test]
385    fn test_display_path_relative() {
386        let working_dir = PathBuf::from("/home/user/project");
387        let saved = PathBuf::from("/home/user/project/images/cat.png");
388        assert_eq!(display_path(&saved, &working_dir), "images/cat.png");
389    }
390
391    #[test]
392    fn test_display_path_outside_working_dir() {
393        let working_dir = PathBuf::from("/home/user/project");
394        let saved = PathBuf::from("/tmp/images/cat.png");
395        assert_eq!(display_path(&saved, &working_dir), "/tmp/images/cat.png");
396    }
397
398    fn args(prompt: &str) -> GenerateImageArgs {
399        serde_json::from_value(json!({ "prompt": prompt })).unwrap()
400    }
401
402    #[test]
403    fn test_extra_params_all_absent_is_null() {
404        assert_eq!(build_extra_params(&args("a cat")), Value::Null);
405    }
406
407    #[test]
408    fn test_extra_params_blank_strings_filtered() {
409        let mut a = args("a cat");
410        a.size = Some("  ".to_string());
411        a.negative_prompt = Some("".to_string());
412        assert_eq!(build_extra_params(&a), Value::Null);
413    }
414
415    #[test]
416    fn test_extra_params_all_present() {
417        let mut a = args("a cat");
418        a.size = Some(" 1696*960 ".to_string());
419        a.negative_prompt = Some("低分辨率".to_string());
420        a.prompt_extend = Some(false);
421        a.seed = Some(42);
422        let extra = build_extra_params(&a);
423        assert_eq!(extra["size"], json!("1696*960"));
424        assert_eq!(extra["negative_prompt"], json!("低分辨率"));
425        assert_eq!(extra["prompt_extend"], json!(false));
426        assert_eq!(extra["seed"], json!(42));
427        assert_eq!(extra.as_object().unwrap().len(), 4);
428    }
429
430    #[test]
431    fn test_args_deserialize_optional_fields() {
432        let a: GenerateImageArgs = serde_json::from_value(json!({
433            "prompt": "a cat",
434            "size": "1280*1280",
435            "seed": 2147483647u64
436        }))
437        .unwrap();
438        assert_eq!(a.size.as_deref(), Some("1280*1280"));
439        assert_eq!(a.seed, Some(2_147_483_647));
440        assert_eq!(a.negative_prompt, None);
441        assert_eq!(a.prompt_extend, None);
442    }
443}