1use 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
21const MAX_N: u32 = 4;
23
24const 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 #[serde(default)]
38 size: Option<String>,
39 #[serde(default)]
41 negative_prompt: Option<String>,
42 #[serde(default)]
44 prompt_extend: Option<bool>,
45 #[serde(default)]
47 seed: Option<u64>,
48}
49
50fn 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 let n = parsed.n.unwrap_or(1).clamp(1, MAX_N);
187
188 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 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 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 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 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 ToolResult::error(content)
327 } else {
328 ToolResult::success(content)
329 }
330 });
331
332 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
351fn 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
360fn display_path(path: &Path, working_dir: &Path) -> String {
363 if let Ok(rel) = path.strip_prefix(working_dir) {
364 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 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}