| 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 |