返回 CodeWhale
tools.rs
根目录 / crates / tui / src / vision / tools.rs
1 //! `image_analyze` tool — analyze images using a dedicated vision model.
2
3 use std::path::{Component, Path, PathBuf};
4 use std::time::Duration;
5
6 use async_trait::async_trait;
7 use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
8 use serde_json::{Value, json};
9
10 use crate::client::CodewhaleClient;
11 use crate::config::ProviderKind;
12 use crate::config::VisionModelConfig;
13 use crate::llm_client::{LlmError, RetryConfig, sanitize_http_error_body, with_retry};
14 use crate::tools::spec::{
15 ToolCapability, ToolContext, ToolError, ToolResult, ToolSpec, required_str,
16 };
17
18 pub struct ImageAnalyzeTool {
19 config: VisionModelConfig,
20 client: reqwest::Client,
21 route_client: Option<CodewhaleClient>,
22 }
23
24 /// Total envelope for one image_analyze call, retry attempts and response
25 /// body consumption included. reqwest's `read_timeout` is *not* a per-read
26 /// idle bound for the request phase: its timer starts at `send()` and is
27 /// never reset until the response headers arrive, so it silently acts as a
28 /// total deadline on the multi-MB upload plus the full non-streaming vision
29 /// generation — precisely the healthy work a 120s cap used to kill. The
30 /// client therefore bounds only the connect handshake, and this envelope
31 /// (the same 30-minute wall clock as engine streaming,
32 /// `STREAM_MAX_DURATION_SECS`) is the sole total bound; a stalled
33 /// connection errors out through it instead of hanging.
34 const VISION_REQUEST_ENVELOPE: Duration = Duration::from_secs(1800);
35
36 fn vision_request_envelope() -> Duration {
37 if cfg!(test) {
38 Duration::from_secs(2)
39 } else {
40 VISION_REQUEST_ENVELOPE
41 }
42 }
43
44 impl ImageAnalyzeTool {
45 #[cfg(test)]
46 #[must_use]
47 pub fn new(config: VisionModelConfig) -> Self {
48 Self::new_with_route_client(config, None)
49 }
50
51 #[must_use]
52 pub fn new_with_route_client(
53 config: VisionModelConfig,
54 route_client: Option<CodewhaleClient>,
55 ) -> Self {
56 let client = crate::tls::reqwest_client_builder()
57 // Bound only the connect handshake. A client- or request-level
58 // `read_timeout` would start counting at `send()` and never
59 // reset before the response headers, quietly re-introducing a
60 // total deadline on the upload + long non-streaming generation;
61 // the total bound lives in VISION_REQUEST_ENVELOPE instead.
62 .connect_timeout(Duration::from_secs(10))
63 .build()
64 .expect("Failed to build HTTP client");
65 Self {
66 config,
67 client,
68 route_client,
69 }
70 }
71
72 async fn read_image_file(path: &Path) -> Result<(String, String), ToolError> {
73 let bytes = tokio::fs::read(path)
74 .await
75 .map_err(|e| ToolError::execution_failed(format!("Failed to read image file: {e}")))?;
76
77 let mime_type = Self::detect_mime_type(path)?;
78 let base64_data = BASE64.encode(&bytes);
79 Ok((base64_data, mime_type))
80 }
81
82 fn resolve_image_path(workspace: &Path, image_path: &str) -> Result<PathBuf, ToolError> {
83 let image_path_buf = Path::new(image_path);
84 if image_path_buf.components().any(|c| {
85 matches!(
86 c,
87 Component::Prefix(_) | Component::RootDir | Component::ParentDir
88 )
89 }) {
90 return Err(ToolError::execution_failed(
91 "image_path must be a relative path within the workspace and cannot escape it.",
92 ));
93 }
94
95 let workspace = workspace.canonicalize().map_err(|e| {
96 ToolError::execution_failed(format!("Failed to resolve workspace path: {e}"))
97 })?;
98 let candidate = workspace.join(image_path_buf);
99 let resolved = candidate.canonicalize().map_err(|e| {
100 ToolError::execution_failed(format!("Failed to resolve image file: {e}"))
101 })?;
102 if !resolved.starts_with(&workspace) {
103 return Err(ToolError::execution_failed(
104 "image_path must resolve within the workspace and cannot escape it.",
105 ));
106 }
107 Ok(resolved)
108 }
109
110 fn detect_mime_type(path: &Path) -> Result<String, ToolError> {
111 let extension = path
112 .extension()
113 .and_then(|e| e.to_str())
114 .unwrap_or("")
115 .to_lowercase();
116
117 match extension.as_str() {
118 "png" => Ok("image/png".to_string()),
119 "jpg" | "jpeg" => Ok("image/jpeg".to_string()),
120 "gif" => Ok("image/gif".to_string()),
121 "webp" => Ok("image/webp".to_string()),
122 "bmp" => Ok("image/bmp".to_string()),
123 _ => Err(ToolError::execution_failed(format!(
124 "Unsupported image format: {extension}"
125 ))),
126 }
127 }
128
129 fn base_url(&self) -> String {
130 self.config
131 .base_url
132 .clone()
133 .unwrap_or_else(|| "https://api.openai.com/v1".to_string())
134 }
135
136 fn api_key(&self) -> String {
137 self.config.api_key.clone().unwrap_or_default()
138 }
139
140 fn is_xiaomi_mimo_model(model: &str) -> bool {
141 let normalized = model.trim().to_ascii_lowercase();
142 let normalized = normalized.strip_prefix("xiaomi/").unwrap_or(&normalized);
143 normalized.starts_with("mimo-")
144 }
145
146 fn uses_max_completion_tokens(config: &VisionModelConfig) -> bool {
147 if Self::is_xiaomi_mimo_model(&config.model) {
148 return true;
149 }
150
151 let base_url = config.base_url.as_deref().unwrap_or_default();
152 let Ok(url) = reqwest::Url::parse(base_url) else {
153 return false;
154 };
155 let Some(domain) = url.domain() else {
156 return false;
157 };
158
159 domain.eq_ignore_ascii_case("xiaomimimo.com")
160 || domain.to_ascii_lowercase().ends_with(".xiaomimimo.com")
161 }
162
163 fn request_payload(&self, prompt: &str, image_data: &str, mime_type: &str) -> Value {
164 let mut payload = json!({
165 "model": self.config.model,
166 "messages": [
167 {
168 "role": "user",
169 "content": [
170 {"type": "text", "text": prompt},
171 {
172 "type": "image_url",
173 "image_url": {
174 "url": format!("data:{};base64,{}", mime_type, image_data)
175 }
176 }
177 ]
178 }
179 ]
180 });
181
182 let token_limit_field = if Self::uses_max_completion_tokens(&self.config) {
183 "max_completion_tokens"
184 } else {
185 "max_tokens"
186 };
187 let configured_base = self.base_url();
188 let route_cap = self
189 .route_client
190 .as_ref()
191 .filter(|client| {
192 client.base_url().trim_end_matches('/') == configured_base.trim_end_matches('/')
193 })
194 .map_or_else(
195 || {
196 // A standalone `[vision_model]` route has no resolved
197 // max-model-len fact. Do not guess one or let a process
198 // override turn a capability maximum into an unbounded
199 // request; a matched active client above carries exact
200 // route limits when the vision route is shared.
201 crate::route_budget::effective_max_output_tokens_for_route(
202 ProviderKind::Custom,
203 &self.config.model,
204 None,
205 )
206 .min(65_536)
207 },
208 |client| client.effective_max_output_tokens(&self.config.model),
209 );
210 payload[token_limit_field] = json!(route_cap);
211 if let Some(client) = self.route_client.as_ref().filter(|client| {
212 client.base_url().trim_end_matches('/') == configured_base.trim_end_matches('/')
213 }) {
214 client.apply_provider_routing(&mut payload);
215 }
216
217 payload
218 }
219 }
220
221 #[async_trait]
222 impl ToolSpec for ImageAnalyzeTool {
223 fn name(&self) -> &str {
224 "image_analyze"
225 }
226
227 fn description(&self) -> &str {
228 "Analyze an image using the configured vision model. \
229 Supports PNG, JPEG, GIF, WebP, and BMP formats."
230 }
231
232 fn input_schema(&self) -> Value {
233 json!({
234 "type": "object",
235 "properties": {
236 "image_path": {
237 "type": "string",
238 "description": "Path to the image file to analyze"
239 },
240 "prompt": {
241 "type": "string",
242 "description": "Optional prompt to guide the analysis."
243 }
244 },
245 "required": ["image_path"]
246 })
247 }
248
249 fn capabilities(&self) -> Vec<ToolCapability> {
250 vec![ToolCapability::ReadOnly]
251 }
252
253 async fn execute(&self, input: Value, context: &ToolContext) -> Result<ToolResult, ToolError> {
254 let image_path = required_str(&input, "image_path")?;
255 let prompt = input
256 .get("prompt")
257 .and_then(|v| v.as_str())
258 .unwrap_or("Describe this image in detail.");
259
260 let resolved_path = Self::resolve_image_path(&context.workspace, image_path)?;
261 let (image_data, mime_type) = Self::read_image_file(&resolved_path).await?;
262
263 let payload = self.request_payload(prompt, &image_data, &mime_type);
264
265 let url = format!("{}/chat/completions", self.base_url());
266 let api_key = self.api_key();
267
268 let retry_config = RetryConfig {
269 max_retries: 3,
270 initial_delay: 1.0,
271 max_delay: 30.0,
272 enabled: true,
273 ..Default::default()
274 };
275 let _inference = match self.route_client.as_ref() {
276 Some(client) => client.acquire_remote_control_inference_permit().await,
277 None => Some(crate::client::acquire_remote_control_inference_participant().await),
278 };
279
280 let response_json = tokio::time::timeout(vision_request_envelope(), async {
281 let response = with_retry(
282 &retry_config,
283 || {
284 let client = self.client.clone();
285 let url = url.clone();
286 let api_key = api_key.clone();
287 let payload = payload.clone();
288 async move {
289 let response = client
290 .post(&url)
291 .header("Content-Type", "application/json")
292 .header("Authorization", format!("Bearer {api_key}"))
293 .json(&payload)
294 .send()
295 .await
296 .map_err(|e| LlmError::from_reqwest(&e))?;
297
298 let status = response.status();
299 if !status.is_success() {
300 let error_text = response
301 .text()
302 .await
303 .unwrap_or_else(|_| "Unknown error".to_string());
304 let error_text = sanitize_http_error_body(
305 Some("Vision provider"),
306 status.as_u16(),
307 &error_text,
308 );
309 return Err(LlmError::from_http_response(status.as_u16(), &error_text));
310 }
311 Ok(response)
312 }
313 },
314 None,
315 )
316 .await
317 .map_err(|e| ToolError::execution_failed(format!("Vision API request failed: {e}")))?;
318
319 let json: Value = response.json().await.map_err(|e| {
320 ToolError::execution_failed(format!("Failed to parse response: {e}"))
321 })?;
322 Ok::<Value, ToolError>(json)
323 })
324 .await
325 .map_err(|_| ToolError::Timeout {
326 seconds: vision_request_envelope().as_secs(),
327 })??;
328
329 let content = response_json
330 .get("choices")
331 .and_then(|c| c.get(0))
332 .and_then(|c| c.get("message"))
333 .and_then(|m| m.get("content"))
334 .and_then(|c| c.as_str())
335 .unwrap_or("")
336 .to_string();
337
338 let model = response_json
339 .get("model")
340 .and_then(|m| m.as_str())
341 .unwrap_or(&self.config.model)
342 .to_string();
343
344 let result = json!({
345 "analysis": content,
346 "model": model,
347 });
348
349 ToolResult::json(&result)
350 .map_err(|e| ToolError::execution_failed(format!("Failed to serialize result: {e}")))
351 }
352 }
353
354 #[cfg(test)]
355 mod tests {
356 use super::*;
357 use tempfile::tempdir;
358 use wiremock::matchers::{method, path};
359 use wiremock::{Mock, MockServer, ResponseTemplate};
360
361 #[cfg(unix)]
362 fn create_file_symlink(
363 target: &std::path::Path,
364 link: &std::path::Path,
365 ) -> std::io::Result<()> {
366 std::os::unix::fs::symlink(target, link)
367 }
368
369 #[cfg(windows)]
370 fn create_file_symlink(
371 target: &std::path::Path,
372 link: &std::path::Path,
373 ) -> std::io::Result<()> {
374 std::os::windows::fs::symlink_file(target, link)
375 }
376
377 fn fake_config() -> VisionModelConfig {
378 VisionModelConfig {
379 model: "test-vision-model".to_string(),
380 api_key: Some("test-key".to_string()),
381 base_url: Some("https://example.invalid/v1".to_string()),
382 }
383 }
384
385 /// The cap reads the `*_MAX_OUTPUT_TOKENS` override from the process
386 /// environment, as the payload does. Tests that compare the two hold the
387 /// test env lock, so an engine test setting the override cannot land
388 /// between the two reads in a shared process.
389 fn standalone_vision_cap(model: &str) -> u64 {
390 u64::from(
391 crate::route_budget::effective_max_output_tokens_for_route(
392 ProviderKind::Custom,
393 model,
394 None,
395 )
396 .min(65_536),
397 )
398 }
399
400 #[test]
401 fn tool_metadata_is_read_only_and_named_image_analyze() {
402 let tool = ImageAnalyzeTool::new(fake_config());
403 assert_eq!(tool.name(), "image_analyze");
404 assert!(tool.capabilities().contains(&ToolCapability::ReadOnly));
405 }
406
407 #[test]
408 fn mime_type_detection_covers_common_formats() {
409 for (ext, expected) in [
410 ("png", "image/png"),
411 ("PNG", "image/png"),
412 ("jpg", "image/jpeg"),
413 ("jpeg", "image/jpeg"),
414 ("gif", "image/gif"),
415 ("webp", "image/webp"),
416 ("bmp", "image/bmp"),
417 ] {
418 let path = std::path::PathBuf::from(format!("test.{ext}"));
419 let mime = ImageAnalyzeTool::detect_mime_type(&path)
420 .unwrap_or_else(|_| panic!("must detect {ext}"));
421 assert_eq!(mime, expected);
422 }
423 }
424
425 #[test]
426 fn mime_type_detection_rejects_unsupported_extension() {
427 let path = std::path::PathBuf::from("test.svg");
428 let err = ImageAnalyzeTool::detect_mime_type(&path)
429 .expect_err("svg is intentionally out of scope for vision tool");
430 assert!(err.to_string().contains("Unsupported image format"));
431 }
432
433 #[test]
434 fn generic_vision_payload_uses_max_tokens() {
435 let _env = crate::test_support::lock_test_env();
436 let tool = ImageAnalyzeTool::new(fake_config());
437
438 let payload = tool.request_payload("describe", "abc123", "image/png");
439
440 assert_eq!(
441 payload.get("max_tokens").and_then(Value::as_u64),
442 Some(standalone_vision_cap(&tool.config.model))
443 );
444 assert!(payload.get("temperature").is_none());
445 assert!(payload.get("max_completion_tokens").is_none());
446 }
447
448 #[test]
449 fn xiaomi_mimo_vision_payload_uses_max_completion_tokens() {
450 let _env = crate::test_support::lock_test_env();
451 let mut config = fake_config();
452 config.model = "mimo-v2.5".to_string();
453 config.base_url = Some("https://api.xiaomimimo.com/v1".to_string());
454 let tool = ImageAnalyzeTool::new(config);
455
456 let payload = tool.request_payload("describe", "abc123", "image/png");
457
458 assert_eq!(
459 payload.get("max_completion_tokens").and_then(Value::as_u64),
460 Some(standalone_vision_cap(&tool.config.model))
461 );
462 assert!(payload.get("temperature").is_none());
463 assert!(payload.get("max_tokens").is_none());
464 }
465
466 #[test]
467 fn xiaomi_mimo_vision_payload_uses_max_completion_tokens_with_custom_proxy() {
468 let _env = crate::test_support::lock_test_env();
469 let mut config = fake_config();
470 config.model = "mimo-v2.5".to_string();
471 config.base_url = Some("https://vision-proxy.example.invalid/v1".to_string());
472 let tool = ImageAnalyzeTool::new(config);
473
474 let payload = tool.request_payload("describe", "abc123", "image/png");
475
476 assert_eq!(
477 payload.get("max_completion_tokens").and_then(Value::as_u64),
478 Some(standalone_vision_cap(&tool.config.model))
479 );
480 assert!(payload.get("max_tokens").is_none());
481 }
482
483 #[test]
484 fn vision_vendor_pin_requires_the_matching_bound_route() {
485 let _lock = crate::test_support::lock_test_env();
486 let base_url = "http://127.0.0.1:18080/v1";
487 let client = CodewhaleClient::new(&crate::config::Config {
488 provider: Some("openrouter".to_string()),
489 providers: Some(crate::config::ProvidersConfig {
490 openrouter: crate::config::ProviderConfig {
491 api_key: Some("fixture-openrouter-key".to_string()),
492 base_url: Some(base_url.to_string()),
493 model: Some("fixture/vision".to_string()),
494 vendor: Some("chutes/region-fixture".to_string()),
495 ..Default::default()
496 },
497 ..Default::default()
498 }),
499 ..Default::default()
500 })
501 .unwrap();
502 for (vision_base, matched_client, pinned) in [
503 (base_url, Some(client.clone()), true),
504 ("http://127.0.0.1:18081/v1", Some(client.clone()), false),
505 (base_url, None, false),
506 ] {
507 let tool = ImageAnalyzeTool::new_with_route_client(
508 VisionModelConfig {
509 model: "fixture/vision".to_string(),
510 api_key: Some("fixture-vision-key".to_string()),
511 base_url: Some(vision_base.to_string()),
512 },
513 matched_client,
514 );
515 let body = tool.request_payload("describe", "abc123", "image/png");
516 if pinned {
517 assert_eq!(
518 body["provider"],
519 json!({
520 "order": ["chutes/region-fixture"], "allow_fallbacks": false
521 })
522 );
523 } else {
524 assert!(body.get("provider").is_none());
525 }
526 }
527 }
528
529 #[test]
530 fn matched_vision_route_uses_bound_client_window_cap() {
531 let _lock = crate::test_support::lock_test_env();
532 let _canonical =
533 crate::test_support::EnvVarGuard::set("CODEWHALE_MAX_OUTPUT_TOKENS", "384000");
534 let base_url = "http://127.0.0.1:18080/v1".to_string();
535 let model = "DeepSeek-V4-Flash".to_string();
536 let client = CodewhaleClient::new(&crate::config::Config {
537 provider: Some("vllm".to_string()),
538 providers: Some(crate::config::ProvidersConfig {
539 vllm: crate::config::ProviderConfig {
540 base_url: Some(base_url.clone()),
541 model: Some(model.clone()),
542 context_window: Some(327_680),
543 ..crate::config::ProviderConfig::default()
544 },
545 ..crate::config::ProvidersConfig::default()
546 }),
547 ..crate::config::Config::default()
548 })
549 .expect("bound vLLM client");
550 let tool = ImageAnalyzeTool::new_with_route_client(
551 VisionModelConfig {
552 model,
553 api_key: None,
554 base_url: Some(base_url),
555 },
556 Some(client),
557 );
558
559 let payload = tool.request_payload("describe", "abc123", "image/png");
560 assert_eq!(payload["max_tokens"], 325_632);
561 }
562
563 #[tokio::test]
564 async fn execute_rejects_absolute_path() {
565 // Trust-boundary pin: image_path must stay inside the workspace
566 // — an absolute path or a `..`-traversing path must reject
567 // before any base64 / API call.
568 let tmp = tempdir().expect("tempdir");
569 let ctx = ToolContext::new(tmp.path().to_path_buf());
570 let tool = ImageAnalyzeTool::new(fake_config());
571 let outside_workspace = if cfg!(windows) {
572 r"C:\Windows\System32\drivers\etc\hosts"
573 } else {
574 "/etc/hosts"
575 };
576 let err = tool
577 .execute(json!({"image_path": outside_workspace}), &ctx)
578 .await
579 .expect_err("absolute path must reject");
580 assert!(
581 err.to_string()
582 .contains("relative path within the workspace"),
583 "error must call out the workspace boundary; got {err}"
584 );
585 }
586
587 #[tokio::test]
588 async fn execute_rejects_parent_dir_traversal() {
589 let tmp = tempdir().expect("tempdir");
590 let ctx = ToolContext::new(tmp.path().to_path_buf());
591 let tool = ImageAnalyzeTool::new(fake_config());
592 let err = tool
593 .execute(json!({"image_path": "../escape.png"}), &ctx)
594 .await
595 .expect_err("`..`-traversal must reject");
596 assert!(
597 err.to_string()
598 .contains("relative path within the workspace"),
599 "error must call out the workspace boundary; got {err}"
600 );
601 }
602
603 #[tokio::test]
604 async fn execute_rejects_symlink_that_resolves_outside_workspace() {
605 let workspace = tempdir().expect("workspace tempdir");
606 let outside = tempdir().expect("outside tempdir");
607 let outside_image = outside.path().join("outside.png");
608 std::fs::write(&outside_image, b"not a real png").expect("write outside image");
609 let link = workspace.path().join("linked.png");
610 if let Err(err) = create_file_symlink(&outside_image, &link) {
611 eprintln!("skipping symlink assertion: {err}");
612 return;
613 }
614
615 let ctx = ToolContext::new(workspace.path().to_path_buf());
616 let tool = ImageAnalyzeTool::new(fake_config());
617 let err = tool
618 .execute(json!({"image_path": "linked.png"}), &ctx)
619 .await
620 .expect_err("symlink target outside workspace must reject before reading");
621 assert!(
622 err.to_string().contains("resolve within the workspace"),
623 "error must call out the canonical workspace boundary; got {err}"
624 );
625 }
626
627 fn vision_response_body() -> Value {
628 json!({
629 "model": "test-vision-model",
630 "choices": [
631 { "message": { "content": "a red square" } }
632 ]
633 })
634 }
635
636 fn tool_with_base_url(base_url: String) -> ImageAnalyzeTool {
637 ImageAnalyzeTool::new(VisionModelConfig {
638 model: "test-vision-model".to_string(),
639 api_key: Some("test-key".to_string()),
640 base_url: Some(base_url),
641 })
642 }
643
644 fn write_workspace_image(workspace: &std::path::Path) {
645 std::fs::write(workspace.join("sample.png"), b"not a real png")
646 .expect("write sample image");
647 }
648
649 #[tokio::test]
650 async fn envelope_bounds_a_stalled_vision_provider() {
651 let server = MockServer::start().await;
652 // The stalled provider never answers within the test envelope: the
653 // upload + non-streaming generation window must be cut off by the
654 // envelope, not by a client read timeout (which reqwest turns into
655 // a hidden total deadline from `send()`).
656 Mock::given(method("POST"))
657 .and(path("/chat/completions"))
658 .respond_with(
659 ResponseTemplate::new(200)
660 .set_body_json(vision_response_body())
661 .set_delay(Duration::from_secs(30)),
662 )
663 .mount(&server)
664 .await;
665
666 let workspace = tempdir().expect("workspace tempdir");
667 write_workspace_image(workspace.path());
668 let ctx = ToolContext::new(workspace.path().to_path_buf());
669 let tool = tool_with_base_url(server.uri());
670
671 let err = tool
672 .execute(json!({"image_path": "sample.png"}), &ctx)
673 .await
674 .expect_err("a provider that never answers must hit the envelope");
675 assert!(
676 err.to_string().contains("timed out after"),
677 "envelope timeout must be reported as such; got {err}"
678 );
679 }
680
681 #[tokio::test]
682 async fn prompt_answer_within_the_envelope_is_returned() {
683 let server = MockServer::start().await;
684 Mock::given(method("POST"))
685 .and(path("/chat/completions"))
686 .respond_with(ResponseTemplate::new(200).set_body_json(vision_response_body()))
687 .mount(&server)
688 .await;
689
690 let workspace = tempdir().expect("workspace tempdir");
691 write_workspace_image(workspace.path());
692 let ctx = ToolContext::new(workspace.path().to_path_buf());
693 let tool = tool_with_base_url(server.uri());
694
695 let result = tool
696 .execute(
697 json!({"image_path": "sample.png", "prompt": "what is this?"}),
698 &ctx,
699 )
700 .await
701 .expect("a prompt answer must flow through the envelope");
702 let payload: Value =
703 serde_json::from_str(&result.content).expect("tool result must carry json");
704 assert_eq!(payload["analysis"], "a red square");
705 }
706 }
707
707 lines RUST