返回 CodeWhale
tests.rs
根目录 / crates / tui / src / mcp / tests.rs
1 use super::headers::{MCP_HTTP_ACCEPT, is_safe_custom_header, with_default_mcp_http_headers};
2 use super::*;
3 use reqwest::header::{ACCEPT, CONTENT_TYPE};
4 use std::collections::VecDeque;
5 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering as AtomicOrdering};
6 use std::sync::{Arc, Mutex, OnceLock};
7 #[cfg(unix)]
8 use tokio::io::AsyncBufReadExt;
9
10 fn test_http_client() -> reqwest::Client {
11 let _ = rustls::crypto::ring::default_provider().install_default();
12 crate::tls::reqwest_client()
13 }
14
15 async fn lock_mcp_loopback_tests() -> tokio::sync::MutexGuard<'static, ()> {
16 static LOCK: OnceLock<tokio::sync::Mutex<()>> = OnceLock::new();
17 LOCK.get_or_init(|| tokio::sync::Mutex::new(()))
18 .lock()
19 .await
20 }
21
22 struct WorkspaceTrustConfigGuard {
23 config_path: PathBuf,
24 _codewhale_config_path: crate::test_support::EnvVarGuard,
25 _deepseek_config_path: crate::test_support::EnvVarGuard,
26 _env_lock: crate::test_support::TestEnvLock,
27 }
28
29 fn workspace_trust_config_guard(workspace: &Path) -> WorkspaceTrustConfigGuard {
30 let env_lock = crate::test_support::lock_test_env();
31 let config_path = workspace
32 .parent()
33 .unwrap_or(workspace)
34 .join("user-config")
35 .join("config.toml");
36 if let Some(parent) = config_path.parent() {
37 fs::create_dir_all(parent).unwrap();
38 }
39 let codewhale_config_path =
40 crate::test_support::EnvVarGuard::set("CODEWHALE_CONFIG_PATH", config_path.as_os_str());
41 let deepseek_config_path = crate::test_support::EnvVarGuard::remove("DEEPSEEK_CONFIG_PATH");
42
43 WorkspaceTrustConfigGuard {
44 config_path,
45 _codewhale_config_path: codewhale_config_path,
46 _deepseek_config_path: deepseek_config_path,
47 _env_lock: env_lock,
48 }
49 }
50
51 fn write_workspace_trust_config(config_path: &Path, workspace: &Path) {
52 let workspace = workspace
53 .canonicalize()
54 .unwrap_or_else(|_| workspace.to_path_buf());
55 let key = workspace
56 .to_string_lossy()
57 .replace('\\', "\\\\")
58 .replace('"', "\\\"");
59 fs::write(
60 config_path,
61 format!("[projects.\"{key}\"]\ntrust_level = \"trusted\"\n"),
62 )
63 .unwrap();
64 }
65
66 fn mark_workspace_trusted(workspace: &Path) -> WorkspaceTrustConfigGuard {
67 let guard = workspace_trust_config_guard(workspace);
68 write_workspace_trust_config(&guard.config_path, workspace);
69 guard
70 }
71
72 #[test]
73 fn test_mcp_config_defaults() {
74 let config = McpConfig::default();
75 assert_eq!(config.timeouts.connect_timeout, 10);
76 assert_eq!(config.timeouts.execute_timeout, 60);
77 assert_eq!(config.timeouts.read_timeout, 120);
78 assert!(config.servers.is_empty());
79 }
80
81 #[test]
82 fn reviewed_remote_endpoint_identity_normalizes_case_idna_and_default_ports() {
83 let canonical = reviewed_remote_endpoint_identity("https://example.com/mcp").unwrap();
84 assert_eq!(
85 reviewed_remote_endpoint_identity("https://EXAMPLE.COM:443/mcp").unwrap(),
86 canonical
87 );
88 assert_eq!(
89 reviewed_remote_endpoint_identity("https://BÜCHER.example:443/mcp").unwrap(),
90 reviewed_remote_endpoint_identity("https://xn--bcher-kva.example/mcp").unwrap()
91 );
92 assert_ne!(
93 reviewed_remote_endpoint_identity("https://example.com:444/mcp").unwrap(),
94 canonical
95 );
96 assert!(reviewed_remote_endpoint_identity("http://localhost:8080/mcp").is_ok());
97 assert!(reviewed_remote_endpoint_identity("http://127.0.0.1/mcp").is_ok());
98 assert!(reviewed_remote_endpoint_identity("http://[::1]/mcp").is_ok());
99 }
100
101 #[test]
102 fn reviewed_remote_endpoint_identity_rejects_ambiguous_or_secret_bearing_urls() {
103 for endpoint in [
104 "http://example.com/mcp",
105 "ftp://example.com/mcp",
106 "https://user@example.com/mcp",
107 "https://user:secret@example.com/mcp",
108 "https://example.com/mcp?token=secret",
109 "https://example.com/mcp#fragment",
110 ] {
111 let error = reviewed_remote_endpoint_identity(endpoint)
112 .expect_err("unsafe reviewed endpoint must fail closed")
113 .to_string();
114 assert!(
115 !error.contains("secret"),
116 "endpoint error leaked URL material"
117 );
118 }
119 }
120
121 #[test]
122 fn reviewed_plugin_redirects_are_exact_normalized_origin_only() {
123 let approved = reviewed_remote_endpoint_identity("https://BÜCHER.example:443/mcp")
124 .unwrap()
125 .1;
126 let accepted = [
127 "https://xn--bcher-kva.example/next",
128 "https://BÜCHER.example:443/next?cursor=opaque",
129 ];
130 for endpoint in accepted {
131 assert!(reviewed_redirect_matches_origin(
132 &reqwest::Url::parse(endpoint).unwrap(),
133 &approved
134 ));
135 }
136
137 let rejected = [
138 "http://xn--bcher-kva.example/next",
139 "https://user@xn--bcher-kva.example/next",
140 "https://xn--bcher-kva.example:444/next",
141 "https://other.example/next",
142 ];
143 for endpoint in rejected {
144 assert!(!reviewed_redirect_matches_origin(
145 &reqwest::Url::parse(endpoint).unwrap(),
146 &approved
147 ));
148 }
149 }
150
151 #[test]
152 fn reviewed_plugin_remote_proxy_policy_never_reads_ambient_environment() {
153 let reads = std::cell::Cell::new(0_u32);
154 let builder = configure_mcp_proxy(crate::tls::reqwest_client_builder(), true, |_| {
155 reads.set(reads.get() + 1);
156 Ok("http://127.0.0.1:9999".to_string())
157 });
158
159 assert_eq!(
160 reads.get(),
161 0,
162 "reviewed remotes must not read proxy values"
163 );
164 builder
165 .build()
166 .expect("explicit no-proxy client must remain buildable");
167 }
168
169 #[test]
170 fn user_authored_mcp_proxy_policy_keeps_environment_support() {
171 let requested = std::cell::RefCell::new(Vec::new());
172 let builder = configure_mcp_proxy(crate::tls::reqwest_client_builder(), false, |name| {
173 requested.borrow_mut().push(name.to_string());
174 match name {
175 "HTTPS_PROXY" => Ok("http://127.0.0.1:8080".to_string()),
176 _ => Err(std::env::VarError::NotPresent),
177 }
178 });
179
180 assert_eq!(
181 requested.into_inner(),
182 vec![
183 "HTTPS_PROXY".to_string(),
184 "NO_PROXY".to_string(),
185 "no_proxy".to_string(),
186 ]
187 );
188 builder
189 .build()
190 .expect("user-authored proxy client must remain buildable");
191 }
192
193 #[test]
194 fn test_mcp_config_parse() {
195 let json = r#"{
196 "timeouts": {
197 "connect_timeout": 15,
198 "execute_timeout": 90
199 },
200 "servers": {
201 "test": {
202 "command": "node",
203 "args": ["server.js"],
204 "env": {"FOO": "bar"}
205 }
206 }
207 }"#;
208
209 let config: McpConfig = serde_json::from_str(json).unwrap();
210 assert_eq!(config.timeouts.connect_timeout, 15);
211 assert_eq!(config.timeouts.execute_timeout, 90);
212 assert_eq!(config.timeouts.read_timeout, 120); // default
213 assert!(config.servers.contains_key("test"));
214
215 let server = config.servers.get("test").unwrap();
216 assert_eq!(server.command, Some("node".to_string()));
217 assert_eq!(server.args, vec!["server.js"]);
218 assert_eq!(server.env.get("FOO"), Some(&"bar".to_string()));
219 }
220
221 #[test]
222 fn mcp_pool_parse_prefixed_name_rejects_ambiguous_configured_server_prefixes() {
223 let config: McpConfig = serde_json::from_str(
224 r#"{
225 "servers": {
226 "my": {"command": "node"},
227 "my_db": {"command": "node"}
228 }
229 }"#,
230 )
231 .unwrap();
232 let pool = McpPool::new(config);
233
234 let error = pool
235 .parse_prefixed_name("mcp_my_db_execute_sql")
236 .expect_err("configured server-prefix collisions must fail closed");
237 assert!(error.to_string().contains("Unknown MCP tool name"));
238 }
239
240 #[test]
241 fn mcp_server_config_parses_custom_headers() {
242 let json = r#"{
243 "servers": {
244 "hf": {
245 "url": "https://example.invalid/mcp",
246 "headers": {
247 "Authorization": "Bearer tok",
248 "X-Org": "anthropic"
249 }
250 }
251 }
252 }"#;
253 let cfg: McpConfig = serde_json::from_str(json).unwrap();
254 let hf = cfg.servers.get("hf").expect("server present");
255 assert_eq!(
256 hf.headers.get("Authorization"),
257 Some(&"Bearer tok".to_string())
258 );
259 assert_eq!(hf.headers.get("X-Org"), Some(&"anthropic".to_string()));
260 }
261
262 #[test]
263 fn mcp_server_config_parses_remote_auth_fields() {
264 let json = r#"{
265 "servers": {
266 "remote": {
267 "url": "https://example.invalid/mcp",
268 "env_http_headers": {
269 "X-Api-Key": "REMOTE_MCP_KEY"
270 },
271 "bearer_token_env_var": "REMOTE_MCP_TOKEN",
272 "scopes": ["tools/read", "tools/write"],
273 "oauth": {
274 "client_id": "client-123"
275 },
276 "oauth_resource": "https://example.invalid"
277 }
278 }
279 }"#;
280 let cfg: McpConfig = serde_json::from_str(json).unwrap();
281 let remote = cfg.servers.get("remote").expect("server present");
282 assert_eq!(
283 remote.env_headers.get("X-Api-Key"),
284 Some(&"REMOTE_MCP_KEY".to_string())
285 );
286 assert_eq!(
287 remote.bearer_token_env_var.as_deref(),
288 Some("REMOTE_MCP_TOKEN")
289 );
290 assert_eq!(remote.scopes, vec!["tools/read", "tools/write"]);
291 assert_eq!(remote.oauth_client_id(), Some("client-123"));
292 assert_eq!(
293 remote.oauth_resource.as_deref(),
294 Some("https://example.invalid")
295 );
296 }
297
298 #[test]
299 fn mcp_server_config_omits_headers_when_empty() {
300 // Empty headers map should not appear in the serialized output —
301 // older mcp.json files written before v0.8.31 must round-trip
302 // unchanged so a `mcp save` from a fresh install doesn't add
303 // dead keys.
304 let cfg = McpServerConfig {
305 command: Some("node".into()),
306 args: vec!["server.js".into()],
307 env: HashMap::new(),
308 cwd: None,
309 url: None,
310 transport: None,
311 connect_timeout: None,
312 execute_timeout: None,
313 read_timeout: None,
314 disabled: false,
315 enabled: true,
316 required: false,
317 enabled_tools: Vec::new(),
318 disabled_tools: Vec::new(),
319 headers: HashMap::new(),
320 env_headers: HashMap::new(),
321 bearer_token_env_var: None,
322 scopes: Vec::new(),
323 oauth: None,
324 oauth_resource: None,
325 reviewed_plugin: None,
326 };
327 let serialized = serde_json::to_string(&cfg).unwrap();
328 assert!(
329 !serialized.contains("\"headers\""),
330 "empty headers must be omitted: {serialized}"
331 );
332 assert!(
333 !serialized.contains("\"env_headers\""),
334 "empty env_headers must be omitted: {serialized}"
335 );
336 assert!(
337 !serialized.contains("\"scopes\""),
338 "empty scopes must be omitted: {serialized}"
339 );
340 assert!(
341 !serialized.contains("\"oauth\""),
342 "empty oauth config must be omitted: {serialized}"
343 );
344 }
345
346 #[test]
347 fn expand_env_placeholders_expands_value_from_environment() {
348 let _lock = crate::test_support::lock_test_env();
349 let _secret =
350 crate::test_support::EnvVarGuard::set("MCP_TEST_SECRET_TOKEN", "test-secret-123456");
351 let mut env = HashMap::new();
352 env.insert(
353 "API_TOKEN".to_string(),
354 "${MCP_TEST_SECRET_TOKEN}".to_string(),
355 );
356
357 let expanded = expand_env_placeholders_map(&env, "env").unwrap();
358
359 assert_eq!(
360 expanded.get("API_TOKEN").map(String::as_str),
361 Some("test-secret-123456")
362 );
363 }
364
365 #[test]
366 fn expand_env_placeholders_reports_missing_variable_without_secret_value() {
367 let _lock = crate::test_support::lock_test_env();
368 let _missing = crate::test_support::EnvVarGuard::remove("MCP_TEST_MISSING_SECRET");
369
370 let err = expand_env_placeholders("Bearer ${MCP_TEST_MISSING_SECRET}")
371 .expect_err("missing env should fail")
372 .to_string();
373
374 // The error must name the variable but must not leak the surrounding
375 // value (which in practice carries the secret).
376 assert!(err.contains("MCP_TEST_MISSING_SECRET"));
377 assert!(!err.contains("Bearer "));
378 }
379
380 #[test]
381 fn reviewed_plugin_environment_uses_only_the_pre_dotenv_snapshot() {
382 let _lock = crate::test_support::lock_test_env();
383 let dir = tempfile::tempdir().unwrap();
384 let plugin_base = dir.path().join("plugins/env-snapshot");
385 fs::create_dir_all(&plugin_base).unwrap();
386 fs::write(
387 plugin_base.join("plugin.toml"),
388 "schema_version = 1\n[plugin]\nname = \"env-snapshot\"\nversion = \"1.0.0\"\n",
389 )
390 .unwrap();
391 let (_, authority) = active_plugin_fixture(&plugin_base);
392 let snapshot = crate::plugins::HostEnvironment::from_entries([(
393 OsString::from("PLUGIN_SNAPSHOT_TOKEN"),
394 OsString::from("captured-before-dotenv"),
395 )]);
396 let mut server = test_server_config();
397 server
398 .env
399 .insert("TOKEN".to_string(), "${PLUGIN_SNAPSHOT_TOKEN}".to_string());
400 server.reviewed_plugin =
401 Some(ReviewedPluginMcpSource::from_authority(authority, None, Arc::new(snapshot)).unwrap());
402 let _late_dotenv = crate::test_support::EnvVarGuard::set(
403 "PLUGIN_SNAPSHOT_TOKEN",
404 "workspace-dotenv-must-not-win",
405 );
406
407 let expanded = expanded_mcp_stdio_env(&server).unwrap();
408 assert_eq!(expanded["TOKEN"], "captured-before-dotenv");
409
410 server.reviewed_plugin.as_mut().unwrap().host_environment =
411 Arc::new(crate::plugins::HostEnvironment::from_entries([]));
412 let error = expanded_mcp_stdio_env(&server)
413 .expect_err("a value present only after dotenv must fail closed");
414 assert!(
415 format!("{error:#}").contains("PLUGIN_SNAPSHOT_TOKEN"),
416 "unexpected missing-snapshot error: {error:#}"
417 );
418 assert!(!format!("{error:#}").contains("workspace-dotenv-must-not-win"));
419 }
420
421 fn write_path_only_test_command(dir: &Path) -> String {
422 let command = "codewhale-mcp-path-only-test";
423 #[cfg(windows)]
424 let file_name = format!("{command}.exe");
425 #[cfg(not(windows))]
426 let file_name = command.to_string();
427 let path = dir.join(file_name);
428 fs::write(&path, b"test executable").expect("write path-only test command");
429 #[cfg(unix)]
430 {
431 use std::os::unix::fs::PermissionsExt;
432
433 let mut permissions = fs::metadata(&path)
434 .expect("path-only command metadata")
435 .permissions();
436 permissions.set_mode(0o755);
437 fs::set_permissions(&path, permissions).expect("make path-only test command executable");
438 }
439 command.to_string()
440 }
441
442 #[test]
443 fn static_mcp_command_uses_expanded_sanitized_stdio_path() {
444 let _lock = crate::test_support::lock_test_env();
445 let temp = tempfile::tempdir().expect("tempdir");
446 let command = write_path_only_test_command(temp.path());
447 let _path = crate::test_support::EnvVarGuard::set(
448 "CODEWHALE_MCP_PATH_ONLY_DIR",
449 temp.path().as_os_str(),
450 );
451 let _secret = crate::test_support::EnvVarGuard::set(
452 "CODEWHALE_MCP_STATIC_TEST_SECRET",
453 "must-not-reach-child",
454 );
455 let mut server = test_server_config();
456 server.command = Some(command);
457 server.env.insert(
458 "PATH".to_string(),
459 "${CODEWHALE_MCP_PATH_ONLY_DIR}".to_string(),
460 );
461
462 assert_eq!(
463 static_mcp_command_availability(&server).expect("static command check"),
464 McpCommandAvailability::Available
465 );
466
467 let child_env = mcp_stdio_child_env(&server).expect("stdio child env");
468 assert_eq!(
469 env_value(&child_env, "PATH"),
470 Some(temp.path().as_os_str()),
471 "expanded server PATH must override the inherited PATH"
472 );
473 assert!(
474 child_env
475 .iter()
476 .all(|(key, _)| key != "CODEWHALE_MCP_STATIC_TEST_SECRET"),
477 "static lookup must use the same sanitized parent environment as spawn"
478 );
479
480 let expanded_env = expand_env_placeholders_map(&server.env, "env").expect("expanded env");
481 let mut old_spawn_command = tokio::process::Command::new("unused-test-command");
482 crate::child_env::apply_to_tokio_command_mcp(
483 &mut old_spawn_command,
484 crate::child_env::string_map_env(&expanded_env),
485 );
486 let old_spawn_env = old_spawn_command
487 .as_std()
488 .get_envs()
489 .map(|(key, value)| {
490 (
491 key.to_os_string(),
492 value.expect("spawn env value").to_os_string(),
493 )
494 })
495 .collect::<HashMap<_, _>>();
496 let static_env = child_env.into_iter().collect::<HashMap<_, _>>();
497 assert_eq!(
498 static_env, old_spawn_env,
499 "static lookup and the pre-fix spawn helper must receive identical environments"
500 );
501 }
502
503 #[cfg(not(windows))]
504 #[test]
505 fn static_mcp_command_reports_missing_with_server_path_override() {
506 let temp = tempfile::tempdir().expect("tempdir");
507 let mut server = test_server_config();
508 server.command = Some("codewhale-mcp-command-that-does-not-exist".to_string());
509 server.env.insert(
510 "PATH".to_string(),
511 temp.path().to_string_lossy().into_owned(),
512 );
513
514 assert_eq!(
515 static_mcp_command_availability(&server).expect("static command check"),
516 McpCommandAvailability::Missing
517 );
518 }
519
520 #[test]
521 fn static_mcp_command_reports_invalid_path_expansion() {
522 let _lock = crate::test_support::lock_test_env();
523 let _missing = crate::test_support::EnvVarGuard::remove("CODEWHALE_MCP_MISSING_PATH_DIR");
524 let mut server = test_server_config();
525 server.command = Some("codewhale-mcp-command".to_string());
526 server.env.insert(
527 "PATH".to_string(),
528 "do-not-leak-${CODEWHALE_MCP_MISSING_PATH_DIR}-also-secret".to_string(),
529 );
530
531 let error = static_mcp_command_availability(&server)
532 .expect_err("missing PATH placeholder must fail static validation");
533 let error = format!("{error:#}");
534 assert!(error.contains("CODEWHALE_MCP_MISSING_PATH_DIR"));
535 assert!(!error.contains("codewhale-mcp-command"));
536 assert!(!error.contains("do-not-leak"));
537 assert!(!error.contains("also-secret"));
538 }
539
540 #[cfg(unix)]
541 fn write_unix_test_command(path: &Path, mode: u32) {
542 use std::os::unix::fs::PermissionsExt;
543
544 fs::write(path, b"#!/bin/sh\nexit 0\n").expect("write Unix test command");
545 let mut permissions = fs::metadata(path)
546 .expect("Unix test command metadata")
547 .permissions();
548 permissions.set_mode(mode);
549 fs::set_permissions(path, permissions).expect("set Unix test command mode");
550 }
551
552 #[cfg(unix)]
553 #[test]
554 fn static_mcp_command_anchors_relative_and_empty_path_to_server_cwd() {
555 let temp = tempfile::tempdir().expect("tempdir");
556 let cwd = temp.path().join("server-cwd");
557 let bin = cwd.join("relative-bin");
558 fs::create_dir_all(&bin).expect("relative bin dir");
559 let relative_command = "codewhale-mcp-relative-path-test";
560 write_unix_test_command(&bin.join(relative_command), 0o755);
561
562 let mut server = test_server_config();
563 server.command = Some(relative_command.to_string());
564 server.cwd = Some(cwd.clone());
565 server
566 .env
567 .insert("PATH".to_string(), "relative-bin".to_string());
568 assert_eq!(
569 static_mcp_command_availability(&server).expect("relative PATH check"),
570 McpCommandAvailability::Available
571 );
572
573 let empty_path_command = "codewhale-mcp-empty-path-test";
574 write_unix_test_command(&cwd.join(empty_path_command), 0o755);
575 server.command = Some(empty_path_command.to_string());
576 server.env.insert("PATH".to_string(), String::new());
577 assert_eq!(
578 static_mcp_command_availability(&server).expect("empty PATH check"),
579 McpCommandAvailability::Available,
580 "an empty Unix PATH entry resolves from the child's cwd"
581 );
582 }
583
584 #[cfg(unix)]
585 #[test]
586 fn static_mcp_command_preserves_literal_name_and_requires_execute_bits() {
587 let temp = tempfile::tempdir().expect("tempdir");
588 let literal_command = " codewhale-mcp-literal-command-test ";
589 write_unix_test_command(&temp.path().join(literal_command), 0o755);
590
591 let mut server = test_server_config();
592 server.command = Some(literal_command.to_string());
593 server.env.insert(
594 "PATH".to_string(),
595 temp.path().to_string_lossy().into_owned(),
596 );
597 assert_eq!(
598 static_mcp_command_availability(&server).expect("literal command check"),
599 McpCommandAvailability::Available,
600 "static validation must not trim the command passed to Command::new"
601 );
602
603 let non_executable = temp.path().join("codewhale-mcp-non-executable-test");
604 write_unix_test_command(&non_executable, 0o644);
605 server.command = Some("codewhale-mcp-non-executable-test".to_string());
606 assert_eq!(
607 static_mcp_command_availability(&server).expect("PATH execute-bit check"),
608 McpCommandAvailability::Missing
609 );
610 server.command = Some(non_executable.to_string_lossy().into_owned());
611 assert_eq!(
612 static_mcp_command_availability(&server).expect("absolute execute-bit check"),
613 McpCommandAvailability::Missing
614 );
615 }
616
617 #[cfg(windows)]
618 #[test]
619 fn static_mcp_command_matches_windows_path_and_extension_rules() {
620 let temp = tempfile::tempdir().expect("tempdir");
621 let command = write_path_only_test_command(temp.path());
622 let mut server = test_server_config();
623 server.command = Some(command);
624 server.env.insert(
625 "Path".to_string(),
626 temp.path().to_string_lossy().into_owned(),
627 );
628
629 assert_eq!(
630 static_mcp_command_availability(&server).expect("case-insensitive PATH check"),
631 McpCommandAvailability::Available
632 );
633
634 server.command = Some(
635 temp.path()
636 .join("codewhale-mcp-path-only-test")
637 .to_string_lossy()
638 .into_owned(),
639 );
640 assert_eq!(
641 static_mcp_command_availability(&server).expect("absolute omitted .exe check"),
642 McpCommandAvailability::Available
643 );
644
645 let pathext_command = "codewhale-mcp-pathext-only-test";
646 fs::write(
647 temp.path().join(format!("{pathext_command}.cmd")),
648 b"@exit /b 0\r\n",
649 )
650 .expect("write PATHEXT-only command");
651 server.command = Some(pathext_command.to_string());
652 server.env.insert("PATHEXT".to_string(), ".CMD".to_string());
653 assert_eq!(
654 static_mcp_command_availability(&server).expect("PATHEXT command check"),
655 McpCommandAvailability::NotChecked,
656 "a child-PATH miss is conservative because Windows still searches implicit fallbacks"
657 );
658 server.command = Some(format!("{pathext_command}.cmd"));
659 assert_eq!(
660 static_mcp_command_availability(&server).expect("explicit .cmd command check"),
661 McpCommandAvailability::Available,
662 "Rust requires non-.exe extensions to be explicit"
663 );
664 }
665
666 #[tokio::test]
667 async fn mcp_http_auth_prefers_static_authorization_over_bearer_env() {
668 let mut headers = HashMap::new();
669 headers.insert("Authorization".to_string(), "Bearer static".to_string());
670 let auth = McpHttpAuth {
671 headers,
672 bearer_token_env_var: Some("PATH".to_string()),
673 ..Default::default()
674 };
675
676 let resolved = auth.resolved_headers().await.unwrap();
677 assert_eq!(
678 resolved.get("Authorization"),
679 Some(&"Bearer static".to_string())
680 );
681 }
682
683 #[tokio::test]
684 async fn mcp_http_auth_uses_bearer_env_when_no_authorization_header() {
685 let auth = McpHttpAuth {
686 bearer_token_env_var: Some("PATH".to_string()),
687 ..Default::default()
688 };
689
690 let resolved = auth.resolved_headers().await.unwrap();
691 assert!(
692 resolved
693 .get("Authorization")
694 .is_some_and(|value| value.starts_with("Bearer ") && value.len() > "Bearer ".len()),
695 "expected PATH-backed bearer header, got {resolved:?}"
696 );
697 }
698
699 #[test]
700 fn is_safe_custom_header_accepts_normal_auth_pairs() {
701 assert!(is_safe_custom_header("Authorization", "Bearer tok"));
702 assert!(is_safe_custom_header("X-Api-Key", "deadbeef"));
703 assert!(is_safe_custom_header("x-org", "anthropic"));
704 }
705
706 #[test]
707 fn is_safe_custom_header_rejects_empty_or_whitespace_key() {
708 assert!(!is_safe_custom_header("", "value"));
709 assert!(!is_safe_custom_header(" ", "value"));
710 }
711
712 #[test]
713 fn is_safe_custom_header_rejects_response_splitting_values() {
714 assert!(
715 !is_safe_custom_header("X-Foo", "abc\r\nSet-Cookie: evil=1"),
716 "CRLF in value must reject — response-splitting defense"
717 );
718 assert!(
719 !is_safe_custom_header("X-Foo", "abc\nbar"),
720 "bare LF in value must reject"
721 );
722 assert!(
723 !is_safe_custom_header("X-Foo", "abc\rbar"),
724 "bare CR in value must reject"
725 );
726 }
727
728 #[test]
729 fn is_safe_custom_header_rejects_protocol_framing_overrides() {
730 // The MCP Streamable HTTP transport relies on its own
731 // Accept / Content-Type values for protocol negotiation;
732 // a stray user override would silently break tool discovery.
733 assert!(!is_safe_custom_header("Accept", "text/plain"));
734 assert!(!is_safe_custom_header("accept", "text/plain"));
735 assert!(!is_safe_custom_header("Content-Type", "text/plain"));
736 assert!(!is_safe_custom_header("CONTENT-TYPE", "x/y"));
737 }
738
739 #[test]
740 fn default_mcp_http_get_accepts_json_and_event_stream() {
741 let client = test_http_client();
742 let request = with_default_mcp_http_headers(client.get("https://example.invalid/mcp"), false)
743 .build()
744 .unwrap();
745 assert_eq!(
746 request.headers().get(ACCEPT).and_then(|v| v.to_str().ok()),
747 Some(MCP_HTTP_ACCEPT)
748 );
749 assert!(
750 request.headers().get(CONTENT_TYPE).is_none(),
751 "SSE GET requests should not advertise a JSON request body"
752 );
753 }
754
755 #[test]
756 fn default_mcp_http_post_accepts_json_and_event_stream() {
757 let client = test_http_client();
758 let request = with_default_mcp_http_headers(client.post("https://example.invalid/mcp"), true)
759 .build()
760 .unwrap();
761 assert_eq!(
762 request.headers().get(ACCEPT).and_then(|v| v.to_str().ok()),
763 Some(MCP_HTTP_ACCEPT)
764 );
765 assert_eq!(
766 request
767 .headers()
768 .get(CONTENT_TYPE)
769 .and_then(|v| v.to_str().ok()),
770 Some("application/json")
771 );
772 }
773
774 #[test]
775 fn streamable_http_transport_stores_headers() {
776 let client = test_http_client();
777 let mut headers = HashMap::new();
778 headers.insert("Authorization".to_string(), "Bearer xyz".to_string());
779 let transport = StreamableHttpTransport::new(
780 client,
781 "https://example.invalid/mcp".to_string(),
782 McpHttpAuth {
783 headers: headers.clone(),
784 ..Default::default()
785 },
786 );
787 assert_eq!(transport.auth.headers, headers);
788 }
789
790 #[test]
791 fn mcp_auth_required_error_item_is_model_visible() {
792 let item = McpPool::mcp_auth_required_error_item("nordic-mcp");
793 assert_eq!(item["error"], "authentication_required");
794 assert_eq!(item["server"], "nordic-mcp");
795 assert!(
796 item["message"]
797 .as_str()
798 .expect("message")
799 .contains("codewhale mcp login nordic-mcp")
800 );
801 }
802
803 #[test]
804 fn test_mcp_config_parse_mcp_servers_alias_and_snapshot() {
805 let dir = tempfile::tempdir().unwrap();
806 let path = dir.path().join("mcp.json");
807 fs::write(
808 &path,
809 r#"{
810 "mcpServers": {
811 "disabled": {
812 "command": "node",
813 "args": ["server.js"],
814 "disabled": true
815 }
816 }
817 }"#,
818 )
819 .unwrap();
820
821 let cfg = load_config(&path).unwrap();
822 assert!(cfg.servers.contains_key("disabled"));
823 let snapshot = manager_snapshot_from_config(&path, true).unwrap();
824 assert!(snapshot.reload_required);
825 assert_eq!(snapshot.servers[0].name, "disabled");
826 assert!(!snapshot.servers[0].enabled);
827 assert_eq!(snapshot.servers[0].error.as_deref(), Some("disabled"));
828 }
829
830 #[test]
831 fn malformed_mcp_config_error_omits_secret_contents_and_keys() {
832 let dir = tempfile::tempdir().unwrap();
833 let path = dir.path().join("mcp.json");
834 let secret = "cw-secret-mcp-config-4507";
835 fs::write(
836 &path,
837 format!(
838 r#"{{"servers":{{"private":{{"headers":{{"Authorization":"{secret}"}} trailing-junk}}}}}}"#
839 ),
840 )
841 .unwrap();
842
843 let error = load_config(&path).expect_err("malformed MCP config must fail");
844 let diagnostic = format!("{error:#}");
845 assert!(!diagnostic.contains(secret), "{diagnostic}");
846 assert!(!diagnostic.contains("Authorization"), "{diagnostic}");
847 assert!(
848 diagnostic.contains("file contents were omitted"),
849 "{diagnostic}"
850 );
851 }
852
853 #[test]
854 fn workspace_mcp_config_merges_with_project_overrides() {
855 let dir = tempfile::tempdir().unwrap();
856 let global_path = dir.path().join("global-mcp.json");
857 let workspace = dir.path().join("workspace");
858 let project_dir = workspace.join(".codewhale");
859 fs::create_dir_all(&project_dir).unwrap();
860 let _trust = mark_workspace_trusted(&workspace);
861 fs::write(
862 &global_path,
863 r#"{
864 "servers": {
865 "global": {"command": "node", "args": ["global.js"]},
866 "shared": {"command": "node", "args": ["global-shared.js"]}
867 }
868 }"#,
869 )
870 .unwrap();
871 fs::write(
872 project_dir.join("mcp.json"),
873 r#"{
874 "servers": {
875 "project": {"command": "php", "args": ["artisan", "boost:mcp"]},
876 "shared": {"command": "php", "args": ["artisan", "shared:mcp"]}
877 }
878 }"#,
879 )
880 .unwrap();
881
882 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
883 let workspace = workspace.canonicalize().unwrap();
884
885 assert!(cfg.servers.contains_key("global"));
886 let project = cfg.servers.get("project").unwrap();
887 assert_eq!(project.command.as_deref(), Some("php"));
888 assert_eq!(project.cwd.as_deref(), Some(workspace.as_path()));
889 let shared = cfg.servers.get("shared").unwrap();
890 assert_eq!(shared.args, vec!["artisan", "shared:mcp"]);
891 assert_eq!(shared.cwd.as_deref(), Some(workspace.as_path()));
892 }
893
894 #[test]
895 fn workspace_manager_snapshot_counts_global_and_project_servers() {
896 let dir = tempfile::tempdir().unwrap();
897 let global_path = dir.path().join("global-mcp.json");
898 let workspace = dir.path().join("workspace");
899 let project_dir = workspace.join(".codewhale");
900 fs::create_dir_all(&project_dir).unwrap();
901 let _trust = mark_workspace_trusted(&workspace);
902 fs::write(
903 &global_path,
904 r#"{
905 "servers": {
906 "chrome-devtools": {"command": "npx", "args": ["-y", "chrome-devtools-mcp@latest"]},
907 "context7": {"command": "npx", "args": ["-y", "@upstash/context7-mcp@latest"]}
908 }
909 }"#,
910 )
911 .unwrap();
912 fs::write(
913 project_dir.join("mcp.json"),
914 r#"{
915 "servers": {
916 "laravel-boost": {"command": "php", "args": ["artisan", "boost:mcp"]}
917 }
918 }"#,
919 )
920 .unwrap();
921
922 let plain = manager_snapshot_from_config(&global_path, false).unwrap();
923 let merged =
924 manager_snapshot_from_config_with_workspace(&global_path, &workspace, false).unwrap();
925
926 assert_eq!(plain.servers.len(), 2);
927 assert_eq!(merged.servers.len(), 3);
928 assert!(
929 merged
930 .servers
931 .iter()
932 .any(|server| server.name == "laravel-boost"),
933 "workspace-aware snapshots must include trusted project MCP servers"
934 );
935 }
936
937 #[test]
938 fn plugin_mcp_servers_are_qualified_and_resolve_relative_cwd() {
939 let dir = tempfile::tempdir().unwrap();
940 let plugin_base = dir.path().join("plugins").join("fleet");
941 fs::create_dir_all(plugin_base.join("servers/local")).unwrap();
942 fs::write(plugin_base.join("servers/local/server.js"), "// server\n").unwrap();
943
944 fs::write(
945 plugin_base.join("plugin.toml"),
946 r#"
947 schema_version = 1
948 [plugin]
949 name = "fleet"
950 version = "1.0.0"
951
952 [mcp_servers.local]
953 command = "node"
954 args = ["server.js"]
955 cwd = "servers/local"
956
957 [mcp_servers.remote]
958 url = "https://example.invalid/mcp"
959
960 [capabilities]
961 network_hosts = ["example.invalid"]
962 "#,
963 )
964 .unwrap();
965 let (plugin, authority) = active_plugin_fixture(&plugin_base);
966 let plugin_for_collision = plugin.clone();
967 let authority_for_collision = authority.clone();
968 let mut config = McpConfig::default();
969 config.servers.insert(
970 "global".to_string(),
971 serde_json::from_str(r#"{"command":"node","args":["global.js"]}"#).unwrap(),
972 );
973
974 let cfg = merge_plugin_mcp_servers_from_plugins(
975 config,
976 vec![("fleet".to_string(), plugin, authority)],
977 )
978 .unwrap();
979
980 assert!(cfg.servers.contains_key("global"));
981
982 let local = cfg.servers.get("plugin-5-fleet-local").unwrap();
983 assert_eq!(local.command.as_deref(), Some("node"));
984 let staged_root = plugin_for_collision.staged_root.as_deref().unwrap();
985 assert_eq!(
986 local.args,
987 vec![
988 staged_root
989 .join("servers/local/server.js")
990 .display()
991 .to_string()
992 ]
993 );
994 assert_eq!(
995 local.cwd.as_deref(),
996 Some(staged_root.join("servers/local").as_path())
997 );
998
999 let remote = cfg.servers.get("plugin-5-fleet-remote").unwrap();
1000 assert_eq!(remote.url.as_deref(), Some("https://example.invalid/mcp"));
1001 assert!(remote.cwd.is_none());
1002
1003 let mut explicit = McpConfig::default();
1004 explicit.servers.insert(
1005 "plugin-5-fleet-local".to_string(),
1006 serde_json::from_str(r#"{"command":"node","args":["explicit.js"]}"#).unwrap(),
1007 );
1008 let collision_safe = merge_plugin_mcp_servers_from_plugins(
1009 explicit,
1010 vec![(
1011 "fleet".to_string(),
1012 plugin_for_collision,
1013 authority_for_collision,
1014 )],
1015 )
1016 .unwrap();
1017 assert_eq!(
1018 collision_safe.servers["plugin-5-fleet-local"].args,
1019 vec!["explicit.js"],
1020 "explicit MCP config must outrank a colliding plugin server"
1021 );
1022 }
1023
1024 #[test]
1025 fn plugin_server_ids_are_unambiguous_across_hyphenated_plugin_and_server_names() {
1026 let left = qualified_plugin_server_name("foo-bar", "baz");
1027 let right = qualified_plugin_server_name("foo", "bar-baz");
1028
1029 assert_eq!(left, "plugin-7-foo-bar-baz");
1030 assert_eq!(right, "plugin-3-foo-bar-baz");
1031 assert_ne!(left, right);
1032 }
1033
1034 #[test]
1035 fn plugin_mcp_adapter_denies_disabled_and_untrusted_bundles() {
1036 let dir = tempfile::tempdir().unwrap();
1037 let plugin_base = dir.path().join("plugin");
1038 fs::create_dir_all(&plugin_base).unwrap();
1039 fs::write(
1040 plugin_base.join("plugin.toml"),
1041 r#"
1042 schema_version = 1
1043 [plugin]
1044 name = "denied"
1045 version = "1.0.0"
1046
1047 [mcp_servers.local]
1048 command = "node"
1049 "#,
1050 )
1051 .unwrap();
1052 let (mut disabled, authority) = active_plugin_fixture(&plugin_base);
1053 disabled.enabled = false;
1054 let mut untrusted = disabled.clone();
1055 untrusted.enabled = true;
1056 untrusted.trust_status = crate::plugins::types::PluginTrustStatus::NeverReviewed;
1057
1058 for plugin in [disabled, untrusted] {
1059 let config = merge_plugin_mcp_servers_from_plugins(
1060 McpConfig::default(),
1061 vec![("denied".to_string(), plugin, authority.clone())],
1062 )
1063 .unwrap();
1064 assert!(
1065 config.servers.is_empty(),
1066 "headless MCP adapter admitted an inactive bundle"
1067 );
1068 }
1069 }
1070
1071 #[test]
1072 fn plugin_mcp_adapter_denies_content_changed_after_snapshot() {
1073 let dir = tempfile::tempdir().unwrap();
1074 let plugin_base = dir.path().join("plugin");
1075 fs::create_dir_all(&plugin_base).unwrap();
1076 let manifest_path = plugin_base.join("plugin.toml");
1077 fs::write(
1078 &manifest_path,
1079 r#"
1080 schema_version = 1
1081 [plugin]
1082 name = "changed"
1083 version = "1.0.0"
1084
1085 [mcp_servers.local]
1086 command = "node"
1087 "#,
1088 )
1089 .unwrap();
1090 let (plugin, authority) = active_plugin_fixture(&plugin_base);
1091 fs::write(plugin_base.join("late-change.txt"), "changed after review").unwrap();
1092
1093 let config = merge_plugin_mcp_servers_from_plugins(
1094 McpConfig::default(),
1095 vec![("changed".to_string(), plugin, authority)],
1096 )
1097 .unwrap();
1098 assert!(config.servers.is_empty());
1099 }
1100
1101 fn registry_with_local_mcp(
1102 name: &str,
1103 base_path: PathBuf,
1104 workspace: &Path,
1105 ) -> crate::plugins::PluginRegistry {
1106 fs::write(
1107 base_path.join("plugin.toml"),
1108 format!(
1109 r#"
1110 schema_version = 1
1111 [plugin]
1112 name = "{name}"
1113 version = "1.0.0"
1114
1115 [mcp_servers.local]
1116 command = "node"
1117 args = ["server.js"]
1118 "#,
1119 ),
1120 )
1121 .unwrap();
1122 let plugins_root = base_path.parent().expect("plugin parent").to_path_buf();
1123 let discovery = crate::plugins::discovery::DiscoveryConfig {
1124 workspace: workspace.to_path_buf(),
1125 user_plugins_dir: plugins_root,
1126 workspace_plugins_dir: workspace.join(".codewhale/plugins-unused"),
1127 builtin_plugin_dirs: Vec::new(),
1128 state_path: workspace
1129 .join("plugin-state")
1130 .join(format!("plugin-state-{name}.json")),
1131 };
1132 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1133 registry.trust(name).unwrap();
1134 registry.enable(name).unwrap();
1135 registry
1136 }
1137
1138 #[test]
1139 fn plugin_mcp_servers_merge_without_project_config() {
1140 let dir = tempfile::tempdir().unwrap();
1141 let global_path = dir.path().join("global-mcp.json");
1142 let workspace = dir.path().join("workspace");
1143 let plugin_base = dir.path().join("plugins").join("fixture");
1144 fs::create_dir_all(&workspace).unwrap();
1145 fs::create_dir_all(&plugin_base).unwrap();
1146 fs::write(
1147 &global_path,
1148 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
1149 )
1150 .unwrap();
1151
1152 let plugins = registry_with_local_mcp("fixture", plugin_base.clone(), &workspace);
1153 let staged_root = plugins
1154 .get("fixture")
1155 .and_then(|plugin| plugin.staged_root.clone())
1156 .expect("trusted plugin should have an immutable runtime snapshot");
1157 let cfg = load_config_with_workspace_and_plugins(&global_path, &workspace, &plugins).unwrap();
1158
1159 assert!(cfg.servers.contains_key("global"));
1160 let qualified_name = qualified_plugin_server_name("fixture", "local");
1161 let local = cfg
1162 .servers
1163 .get(&qualified_name)
1164 .expect("plugin MCP should merge without a project MCP config");
1165 assert_eq!(local.command.as_deref(), Some("node"));
1166 assert_eq!(local.cwd.as_deref(), Some(staged_root.as_path()));
1167 }
1168
1169 #[cfg(unix)]
1170 #[tokio::test]
1171 async fn plugin_mcp_lazy_spawn_denies_component_changed_after_pool_construction() {
1172 use std::os::unix::fs::PermissionsExt;
1173
1174 let dir = tempfile::tempdir().unwrap();
1175 let plugins_root = dir.path().join("plugins");
1176 let plugin_base = plugins_root.join("guarded");
1177 fs::create_dir_all(&plugin_base).unwrap();
1178 let server_path = plugin_base.join("server.sh");
1179 fs::write(&server_path, "#!/bin/sh\nexit 0\n").unwrap();
1180 let mut permissions = fs::metadata(&server_path).unwrap().permissions();
1181 permissions.set_mode(0o700);
1182 fs::set_permissions(&server_path, permissions).unwrap();
1183 fs::write(
1184 plugin_base.join("plugin.toml"),
1185 r#"
1186 schema_version = 1
1187 [plugin]
1188 name = "guarded"
1189 version = "1.0.0"
1190
1191 [mcp_servers.local]
1192 command = "sh"
1193 args = ["server.sh"]
1194 connect_timeout = 1
1195 "#,
1196 )
1197 .unwrap();
1198
1199 let discovery = crate::plugins::discovery::DiscoveryConfig {
1200 workspace: dir.path().join("project"),
1201 user_plugins_dir: plugins_root,
1202 workspace_plugins_dir: dir.path().join("workspace-plugins"),
1203 builtin_plugin_dirs: Vec::new(),
1204 state_path: dir.path().join("plugin-state/state.json"),
1205 };
1206 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1207 registry.trust("guarded").unwrap();
1208 registry.enable("guarded").unwrap();
1209 let active = registry.active_plugins()[0].clone();
1210 let authority = registry.authority_for("guarded").unwrap();
1211 let merged = merge_plugin_mcp_servers_from_plugins(
1212 McpConfig::default(),
1213 vec![("guarded".to_string(), active, authority)],
1214 )
1215 .unwrap();
1216 assert!(
1217 merged.servers["plugin-7-guarded-local"]
1218 .reviewed_plugin
1219 .is_some(),
1220 "plugin provenance must survive through MCP pool construction"
1221 );
1222 let mut pool = McpPool::new(merged);
1223
1224 // Adversarial mutation after trust, enablement, merge, and pool
1225 // construction. If the lazy child executes, it creates this marker before
1226 // closing stdio, so the regression proves denial happened pre-spawn.
1227 let executed_marker = plugin_base.join("executed.marker");
1228 fs::write(&server_path, "#!/bin/sh\n: > executed.marker\nexit 0\n").unwrap();
1229
1230 let error = pool
1231 .get_or_connect("plugin-7-guarded-local")
1232 .await
1233 .err()
1234 .expect("changed reviewed component must be denied before spawn");
1235 let message = format!("{error:#}");
1236 assert!(
1237 message.contains("Refusing to use MCP server 'plugin-7-guarded-local'"),
1238 "unexpected pre-spawn denial: {message}"
1239 );
1240 assert!(message.contains("changed after review"));
1241 assert!(message.contains("/plugin reload"));
1242 assert!(
1243 !executed_marker.exists(),
1244 "mutated MCP component executed despite pre-spawn hash denial"
1245 );
1246 }
1247
1248 #[cfg(unix)]
1249 #[tokio::test]
1250 async fn plugin_mcp_inflight_call_is_cancelled_after_cross_process_revocation() {
1251 let _env_lock = crate::test_support::lock_test_env();
1252 let dir = tempfile::tempdir().unwrap();
1253 let call_marker = dir.path().join("call.marker");
1254 let _call_marker_env = crate::test_support::EnvVarGuard::set(
1255 "CODEWHALE_TEST_PLUGIN_CALL_MARKER",
1256 call_marker.as_os_str(),
1257 );
1258 let plugins_root = dir.path().join("plugins");
1259 let plugin_base = plugins_root.join("revoked");
1260 fs::create_dir_all(&plugin_base).unwrap();
1261 fs::create_dir_all(dir.path().join("project")).unwrap();
1262 fs::write(
1263 plugin_base.join("server.sh"),
1264 r#"#!/bin/sh
1265 trap 'exit 0' TERM INT
1266 while IFS= read -r line; do
1267 case "$line" in
1268 *'"method":"notifications/initialized"'*)
1269 ;;
1270 *'"method":"initialize"'*)
1271 printf '%s\n' '{"jsonrpc":"2.0","id":"1","result":{"protocolVersion":"2024-11-05","serverInfo":{"name":"revocation-test","version":"1.0.0"},"capabilities":{"tools":{}}}}'
1272 ;;
1273 *'"method":"tools/list"'*)
1274 printf '%s\n' '{"jsonrpc":"2.0","id":"2","result":{"tools":[{"name":"wait","description":"Wait until revoked","inputSchema":{"type":"object"}}]}}'
1275 ;;
1276 *'"method":"tools/call"'*)
1277 : > "$CALL_MARKER"
1278 while :; do sleep 1; done
1279 ;;
1280 esac
1281 done
1282 "#,
1283 )
1284 .unwrap();
1285 fs::write(
1286 plugin_base.join("plugin.toml"),
1287 r#"
1288 schema_version = 1
1289 [plugin]
1290 name = "revoked"
1291 version = "1.0.0"
1292
1293 [mcp_servers.local]
1294 command = "sh"
1295 args = ["server.sh"]
1296 connect_timeout = 2
1297 execute_timeout = 30
1298 read_timeout = 30
1299
1300 [mcp_servers.local.env]
1301 CALL_MARKER = "${CODEWHALE_TEST_PLUGIN_CALL_MARKER}"
1302 "#,
1303 )
1304 .unwrap();
1305
1306 let discovery = crate::plugins::discovery::DiscoveryConfig {
1307 workspace: dir.path().join("project"),
1308 user_plugins_dir: plugins_root,
1309 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
1310 builtin_plugin_dirs: Vec::new(),
1311 state_path: dir.path().join("plugin-state/state.json"),
1312 };
1313 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1314 registry.trust("revoked").unwrap();
1315 registry.enable("revoked").unwrap();
1316 let active = registry.active_plugins()[0].clone();
1317 let authority = registry.authority_for("revoked").unwrap();
1318 let merged = merge_plugin_mcp_servers_from_plugins(
1319 McpConfig::default(),
1320 vec![("revoked".to_string(), active, authority)],
1321 )
1322 .unwrap();
1323 let mut pool = McpPool::new(merged);
1324 pool.get_or_connect("plugin-7-revoked-local").await.unwrap();
1325
1326 let call = tokio::spawn(async move {
1327 pool.call_tool("mcp_plugin-7-revoked-local_wait", serde_json::json!({}))
1328 .await
1329 });
1330 for _ in 0..100 {
1331 if call_marker.exists() {
1332 break;
1333 }
1334 if call.is_finished() {
1335 let early = call
1336 .await
1337 .expect("in-flight tool task panicked before reaching the server");
1338 panic!("in-flight tool call ended before reaching the server: {early:?}");
1339 }
1340 tokio::time::sleep(Duration::from_millis(20)).await;
1341 }
1342 assert!(
1343 call_marker.exists(),
1344 "test server never observed the in-flight tool call"
1345 );
1346
1347 let mut external = crate::plugins::discovery::discover_with_config(&discovery);
1348 external.revoke_trust("revoked").unwrap();
1349 let result = tokio::time::timeout(Duration::from_secs(5), call)
1350 .await
1351 .expect("revocation watcher did not cancel the in-flight call")
1352 .unwrap();
1353 let error = result
1354 .expect_err("revoked in-flight call must not complete")
1355 .to_string();
1356 assert!(error.contains("cancelled after authority changed"));
1357 assert!(error.contains("disabled, revoked, or no longer matches"));
1358 }
1359
1360 #[cfg(unix)]
1361 #[tokio::test]
1362 async fn plugin_stdio_authority_cancellation_terminates_an_idle_child() {
1363 let dir = tempfile::tempdir().unwrap();
1364 let plugin_base = dir.path().join("plugins/idle-child");
1365 fs::create_dir_all(&plugin_base).unwrap();
1366 fs::write(
1367 plugin_base.join("plugin.toml"),
1368 "schema_version = 1\n[plugin]\nname = \"idle-child\"\nversion = \"1.0.0\"\n",
1369 )
1370 .unwrap();
1371 let (_, authority) = active_plugin_fixture(&plugin_base);
1372 let mut config = test_server_config();
1373 config.command = Some("sh".to_string());
1374 config.args = vec![
1375 "-c".to_string(),
1376 "trap 'exit 0' TERM INT; while :; do sleep 1; done".to_string(),
1377 ];
1378 config.reviewed_plugin = Some(
1379 ReviewedPluginMcpSource::from_authority(
1380 authority,
1381 None,
1382 Arc::new(crate::plugins::HostEnvironment::capture()),
1383 )
1384 .unwrap(),
1385 );
1386 let cancellation = tokio_util::sync::CancellationToken::new();
1387 let transport = StdioTransport::spawn(
1388 "idle-child",
1389 config.command.as_deref().unwrap(),
1390 &config,
1391 cancellation.clone(),
1392 )
1393 .unwrap();
1394 assert!(transport.child.lock().await.try_wait().unwrap().is_none());
1395
1396 cancellation.cancel();
1397 let deadline = tokio::time::Instant::now() + Duration::from_secs(3);
1398 loop {
1399 if transport.child.lock().await.try_wait().unwrap().is_some() {
1400 break;
1401 }
1402 assert!(
1403 tokio::time::Instant::now() < deadline,
1404 "authority cancellation left the plugin stdio child alive"
1405 );
1406 tokio::time::sleep(Duration::from_millis(20)).await;
1407 }
1408 }
1409
1410 #[cfg(unix)]
1411 #[tokio::test]
1412 async fn plugin_stdio_does_not_surface_reviewed_child_stderr() {
1413 let dir = tempfile::tempdir().unwrap();
1414 let plugin_base = dir.path().join("plugins/stderr-secret");
1415 fs::create_dir_all(&plugin_base).unwrap();
1416 fs::write(
1417 plugin_base.join("plugin.toml"),
1418 "schema_version = 1\n[plugin]\nname = \"stderr-secret\"\nversion = \"1.0.0\"\n",
1419 )
1420 .unwrap();
1421 let (_, authority) = active_plugin_fixture(&plugin_base);
1422 let mut config = test_server_config();
1423 config.command = Some("sh".to_string());
1424 config.args = vec![
1425 "-c".to_string(),
1426 "echo 'ARBITRARY_PLUGIN_CREDENTIAL' 1>&2; exit 1".to_string(),
1427 ];
1428 config.reviewed_plugin = Some(
1429 ReviewedPluginMcpSource::from_authority(
1430 authority,
1431 None,
1432 Arc::new(crate::plugins::HostEnvironment::capture()),
1433 )
1434 .unwrap(),
1435 );
1436 let mut transport = StdioTransport::spawn(
1437 "stderr-secret",
1438 config.command.as_deref().unwrap(),
1439 &config,
1440 tokio_util::sync::CancellationToken::new(),
1441 )
1442 .unwrap();
1443
1444 tokio::time::sleep(Duration::from_millis(100)).await;
1445 let error = transport
1446 .recv()
1447 .await
1448 .expect_err("reviewed child should have closed its transport")
1449 .to_string();
1450 assert!(error.contains("Stdio transport closed"));
1451 assert!(!error.contains("ARBITRARY_PLUGIN_CREDENTIAL"));
1452 }
1453
1454 #[tokio::test]
1455 async fn revoked_plugin_mcp_denies_catalog_tool_resource_and_prompt_operations() {
1456 let dir = tempfile::tempdir().unwrap();
1457 let plugins_root = dir.path().join("plugins");
1458 let plugin_base = plugins_root.join("catalog-guard");
1459 fs::create_dir_all(&plugin_base).unwrap();
1460 fs::create_dir_all(dir.path().join("project")).unwrap();
1461 fs::write(
1462 plugin_base.join("plugin.toml"),
1463 "schema_version = 1\n[plugin]\nname = \"catalog-guard\"\nversion = \"1.0.0\"\n",
1464 )
1465 .unwrap();
1466 let discovery = crate::plugins::discovery::DiscoveryConfig {
1467 workspace: dir.path().join("project"),
1468 user_plugins_dir: plugins_root,
1469 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
1470 builtin_plugin_dirs: Vec::new(),
1471 state_path: dir.path().join("plugin-state/state.json"),
1472 };
1473 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1474 registry.trust("catalog-guard").unwrap();
1475 registry.enable("catalog-guard").unwrap();
1476 let authority = registry.authority_for("catalog-guard").unwrap();
1477
1478 let sent = Arc::new(Mutex::new(Vec::new()));
1479 let mut connection = test_connection(Box::new(ScriptedValueTransport {
1480 sent: Arc::clone(&sent),
1481 responses: VecDeque::new(),
1482 }));
1483 let source = ReviewedPluginMcpSource::from_authority(
1484 authority,
1485 None,
1486 Arc::new(crate::plugins::HostEnvironment::capture()),
1487 )
1488 .unwrap();
1489 connection.config.reviewed_plugin = Some(source.clone());
1490 connection.tools.push(McpTool {
1491 name: "echo".to_string(),
1492 description: None,
1493 input_schema: serde_json::json!({}),
1494 });
1495 connection.resources.push(McpResource {
1496 uri: "memory://one".to_string(),
1497 name: "one".to_string(),
1498 description: None,
1499 mime_type: None,
1500 });
1501 connection.resource_templates.push(McpResourceTemplate {
1502 uri_template: "memory://{id}".to_string(),
1503 name: "memory".to_string(),
1504 description: None,
1505 mime_type: None,
1506 });
1507 connection.prompts.push(McpPrompt {
1508 name: "review".to_string(),
1509 description: None,
1510 arguments: Vec::new(),
1511 });
1512 let mut config = McpConfig::default();
1513 let mut server = test_server_config();
1514 server.reviewed_plugin = Some(source);
1515 config.servers.insert("guarded".to_string(), server);
1516 let mut pool = McpPool::new(config);
1517 pool.connections.insert("guarded".to_string(), connection);
1518 assert_eq!(pool.all_tools().len(), 1);
1519 assert_eq!(pool.all_resources().len(), 1);
1520 assert_eq!(pool.all_resource_templates().len(), 1);
1521 assert_eq!(pool.all_prompts().len(), 1);
1522
1523 let mut external = crate::plugins::discovery::discover_with_config(&discovery);
1524 external.revoke_trust("catalog-guard").unwrap();
1525 assert!(pool.all_tools().is_empty());
1526 assert!(pool.all_resources().is_empty());
1527 assert!(pool.all_resource_templates().is_empty());
1528 assert!(pool.all_prompts().is_empty());
1529
1530 let tool = pool
1531 .call_tool("mcp_guarded_echo", serde_json::json!({}))
1532 .await;
1533 let resource = pool.read_resource("guarded", "memory://one").await;
1534 let prompt = pool
1535 .get_prompt("guarded", "review", serde_json::json!({}))
1536 .await;
1537 let resource_catalog = pool
1538 .call_tool(
1539 "list_mcp_resources",
1540 serde_json::json!({"server": "guarded"}),
1541 )
1542 .await;
1543 let template_catalog = pool
1544 .call_tool(
1545 "list_mcp_resource_templates",
1546 serde_json::json!({"server": "guarded"}),
1547 )
1548 .await;
1549 for result in [tool, resource, prompt, resource_catalog, template_catalog] {
1550 let error = result
1551 .expect_err("revoked plugin MCP operation must fail closed")
1552 .to_string();
1553 assert!(error.contains("Refusing to use MCP server 'guarded'"));
1554 }
1555 assert!(
1556 sent.lock().unwrap().is_empty(),
1557 "revoked plugin MCP operation reached the transport"
1558 );
1559 }
1560
1561 fn cached_reviewed_plugin_catalog_fixture() -> (tempfile::TempDir, PathBuf, PathBuf, McpPool) {
1562 let dir = tempfile::tempdir().unwrap();
1563 let plugins_root = dir.path().join("plugins");
1564 let plugin_base = plugins_root.join("catalog-drift");
1565 fs::create_dir_all(&plugin_base).unwrap();
1566 fs::create_dir_all(dir.path().join("project")).unwrap();
1567 fs::write(
1568 plugin_base.join("plugin.toml"),
1569 "schema_version = 1\n[plugin]\nname = \"catalog-drift\"\nversion = \"1.0.0\"\n",
1570 )
1571 .unwrap();
1572 let discovery = crate::plugins::discovery::DiscoveryConfig {
1573 workspace: dir.path().join("project"),
1574 user_plugins_dir: plugins_root,
1575 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
1576 builtin_plugin_dirs: Vec::new(),
1577 state_path: dir.path().join("plugin-state/state.json"),
1578 };
1579 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1580 registry.trust("catalog-drift").unwrap();
1581 registry.enable("catalog-drift").unwrap();
1582 let authority = registry.authority_for("catalog-drift").unwrap();
1583 let staged_manifest = authority.staged_manifest.clone();
1584
1585 let mut connection = test_connection(Box::new(ScriptedValueTransport {
1586 sent: Arc::new(Mutex::new(Vec::new())),
1587 responses: VecDeque::new(),
1588 }));
1589 let source = ReviewedPluginMcpSource::from_authority(
1590 authority,
1591 None,
1592 Arc::new(crate::plugins::HostEnvironment::capture()),
1593 )
1594 .unwrap();
1595 connection.config.reviewed_plugin = Some(source.clone());
1596 connection.tools.push(McpTool {
1597 name: "echo".to_string(),
1598 description: None,
1599 input_schema: serde_json::json!({}),
1600 });
1601 connection.resources.push(McpResource {
1602 uri: "memory://one".to_string(),
1603 name: "one".to_string(),
1604 description: None,
1605 mime_type: None,
1606 });
1607 connection.resource_templates.push(McpResourceTemplate {
1608 uri_template: "memory://{id}".to_string(),
1609 name: "memory".to_string(),
1610 description: None,
1611 mime_type: None,
1612 });
1613 connection.prompts.push(McpPrompt {
1614 name: "review".to_string(),
1615 description: None,
1616 arguments: Vec::new(),
1617 });
1618 let mut config = McpConfig::default();
1619 let mut server = test_server_config();
1620 server.reviewed_plugin = Some(source);
1621 config.servers.insert("guarded".to_string(), server);
1622 let mut pool = McpPool::new(config);
1623 pool.connections.insert("guarded".to_string(), connection);
1624 assert_eq!(pool.all_tools().len(), 1);
1625 assert_eq!(pool.all_resources().len(), 1);
1626 assert_eq!(pool.all_resource_templates().len(), 1);
1627 assert_eq!(pool.all_prompts().len(), 1);
1628
1629 (dir, plugin_base, staged_manifest, pool)
1630 }
1631
1632 fn assert_reviewed_plugin_catalog_hidden(pool: &McpPool, boundary: &str) {
1633 assert!(pool.all_tools().is_empty());
1634 assert!(pool.all_resources().is_empty());
1635 assert!(pool.all_resource_templates().is_empty());
1636 assert!(pool.all_prompts().is_empty());
1637 assert!(
1638 pool.to_api_tools()
1639 .iter()
1640 .all(|tool| tool.name != "mcp_guarded_echo"),
1641 "{boundary} drift must remove cached reviewed tools from the model API catalog"
1642 );
1643 assert!(pool.parse_prefixed_name("mcp_guarded_echo").is_err());
1644 }
1645
1646 #[test]
1647 fn reviewed_plugin_source_drift_hides_every_cached_catalog_surface() {
1648 let (_dir, plugin_base, _staged_manifest, pool) = cached_reviewed_plugin_catalog_fixture();
1649
1650 fs::write(plugin_base.join("unreviewed-companion.txt"), b"drift").unwrap();
1651
1652 assert_reviewed_plugin_catalog_hidden(&pool, "source");
1653 }
1654
1655 #[cfg(unix)]
1656 #[test]
1657 fn reviewed_plugin_stage_drift_hides_every_cached_catalog_surface() {
1658 use std::io::Write as _;
1659 use std::os::unix::fs::PermissionsExt as _;
1660
1661 let (_dir, _plugin_base, staged_manifest, pool) = cached_reviewed_plugin_catalog_fixture();
1662 std::fs::set_permissions(&staged_manifest, std::fs::Permissions::from_mode(0o600)).unwrap();
1663 std::fs::OpenOptions::new()
1664 .append(true)
1665 .open(&staged_manifest)
1666 .unwrap()
1667 .write_all(b"\n# test-only staged drift\n")
1668 .unwrap();
1669
1670 assert_reviewed_plugin_catalog_hidden(&pool, "staged-tree");
1671 }
1672
1673 #[tokio::test]
1674 async fn reviewed_plugin_oauth_is_disabled_without_network_or_token_mutation() {
1675 let dir = tempfile::tempdir().unwrap();
1676 let plugin_base = dir.path().join("plugins/oauth-disabled");
1677 fs::create_dir_all(&plugin_base).unwrap();
1678 fs::write(
1679 plugin_base.join("plugin.toml"),
1680 "schema_version = 1\n[plugin]\nname = \"oauth-disabled\"\nversion = \"1.0.0\"\n",
1681 )
1682 .unwrap();
1683 let (_, authority) = active_plugin_fixture(&plugin_base);
1684
1685 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1686 let endpoint = format!("http://{}/mcp", listener.local_addr().unwrap());
1687 let mut server = test_server_config();
1688 server.command = None;
1689 server.url = Some(endpoint.clone());
1690 server.reviewed_plugin = Some(
1691 ReviewedPluginMcpSource::from_authority(
1692 authority,
1693 Some(&endpoint),
1694 Arc::new(crate::plugins::HostEnvironment::default()),
1695 )
1696 .unwrap(),
1697 );
1698
1699 assert_eq!(
1700 oauth::auth_status_for_server("plugin-oauth", &server).await,
1701 oauth::McpAuthStatus::Unsupported
1702 );
1703 assert!(oauth::oauth_login_support(&server).await.unwrap().is_none());
1704 assert!(
1705 oauth::McpOAuthRuntime::from_server_config(
1706 "plugin-oauth",
1707 &server,
1708 reqwest::header::HeaderMap::new(),
1709 )
1710 .await
1711 .unwrap()
1712 .is_none()
1713 );
1714 let login_error =
1715 oauth::perform_oauth_login_for_server("plugin-oauth", &server, None, None, None)
1716 .await
1717 .expect_err("plugin OAuth login must be disabled")
1718 .to_string();
1719 assert!(login_error.contains("disabled for plugin-contributed MCP servers"));
1720 let logout_error = oauth::delete_oauth_tokens_for_server("plugin-oauth", &server)
1721 .expect_err("plugin OAuth logout must not touch token storage")
1722 .to_string();
1723 assert!(logout_error.contains("storage is disabled"));
1724
1725 assert!(
1726 tokio::time::timeout(Duration::from_millis(50), listener.accept())
1727 .await
1728 .is_err(),
1729 "plugin OAuth disabled paths must not probe the network"
1730 );
1731 }
1732
1733 fn active_plugin_fixture(
1734 plugin_base: &Path,
1735 ) -> (
1736 crate::plugins::types::LoadedPlugin,
1737 crate::plugins::types::PluginAuthority,
1738 ) {
1739 let plugins_root = plugin_base.parent().expect("plugin parent").to_path_buf();
1740 let root = plugins_root.parent().unwrap_or(&plugins_root).to_path_buf();
1741 let discovery = crate::plugins::discovery::DiscoveryConfig {
1742 workspace: root.join("project"),
1743 user_plugins_dir: plugins_root,
1744 workspace_plugins_dir: root.join("workspace-plugins-unused"),
1745 builtin_plugin_dirs: Vec::new(),
1746 state_path: root.join("plugin-state").join(format!(
1747 "plugin-state-{}.json",
1748 plugin_base.file_name().unwrap().to_string_lossy()
1749 )),
1750 };
1751 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1752 let name = registry
1753 .list()
1754 .first()
1755 .expect("discovered plugin")
1756 .name()
1757 .to_string();
1758 registry.trust(&name).unwrap();
1759 registry.enable(&name).unwrap();
1760 (
1761 registry.get(&name).unwrap().clone(),
1762 registry.authority_for(&name).unwrap(),
1763 )
1764 }
1765
1766 #[test]
1767 fn workspace_mcp_config_ignores_project_file_until_workspace_trusted() {
1768 let dir = tempfile::tempdir().unwrap();
1769 let global_path = dir.path().join("global-mcp.json");
1770 let workspace = dir.path().join("workspace");
1771 let project_dir = workspace.join(".codewhale");
1772 let plugin_base = dir.path().join("plugins").join("fixture");
1773 fs::create_dir_all(&project_dir).unwrap();
1774 fs::create_dir_all(&plugin_base).unwrap();
1775 fs::write(
1776 &global_path,
1777 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
1778 )
1779 .unwrap();
1780 fs::write(
1781 project_dir.join("mcp.json"),
1782 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
1783 )
1784 .unwrap();
1785
1786 let plugins = registry_with_local_mcp("fixture", plugin_base, &workspace);
1787 let cfg = load_config_with_workspace_and_plugins(&global_path, &workspace, &plugins).unwrap();
1788
1789 assert!(cfg.servers.contains_key("global"));
1790 assert!(!cfg.servers.contains_key("project"));
1791 assert!(
1792 cfg.servers
1793 .contains_key(&qualified_plugin_server_name("fixture", "local")),
1794 "user plugin MCP should not be gated by project workspace trust"
1795 );
1796 }
1797
1798 #[test]
1799 fn workspace_mcp_config_ignores_project_local_legacy_trust_marker() {
1800 let dir = tempfile::tempdir().unwrap();
1801 let global_path = dir.path().join("global-mcp.json");
1802 let workspace = dir.path().join("workspace");
1803 let project_dir = workspace.join(".codewhale");
1804 fs::create_dir_all(&project_dir).unwrap();
1805 fs::create_dir_all(workspace.join(".deepseek")).unwrap();
1806 fs::write(workspace.join(".deepseek").join("trusted"), "").unwrap();
1807 fs::write(
1808 &global_path,
1809 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
1810 )
1811 .unwrap();
1812 fs::write(
1813 project_dir.join("mcp.json"),
1814 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
1815 )
1816 .unwrap();
1817
1818 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
1819
1820 assert!(cfg.servers.contains_key("global"));
1821 assert!(!cfg.servers.contains_key("project"));
1822 }
1823
1824 #[test]
1825 fn workspace_mcp_config_ignores_invalid_untrusted_project_file() {
1826 let dir = tempfile::tempdir().unwrap();
1827 let global_path = dir.path().join("global-mcp.json");
1828 let workspace = dir.path().join("workspace");
1829 let project_dir = workspace.join(".codewhale");
1830 fs::create_dir_all(&project_dir).unwrap();
1831 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
1832 fs::write(project_dir.join("mcp.json"), "{ not json").unwrap();
1833
1834 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
1835
1836 assert!(cfg.servers.is_empty());
1837 }
1838
1839 #[test]
1840 fn workspace_mcp_config_rejects_parent_components() {
1841 let dir = tempfile::tempdir().unwrap();
1842 let global_path = dir.path().join("global-mcp.json");
1843 let workspace = dir.path().join("workspace");
1844 let project_dir = workspace.join(".codewhale");
1845 fs::create_dir_all(&project_dir).unwrap();
1846 let _trust = mark_workspace_trusted(&workspace);
1847 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
1848 fs::write(
1849 project_dir.join("mcp.json"),
1850 r#"{"servers": {"project": {"command": "node", "args": ["server.js"]}}}"#,
1851 )
1852 .unwrap();
1853
1854 let workspace_with_parent = workspace.join("..").join("workspace");
1855 let err = load_config_with_workspace(&global_path, &workspace_with_parent)
1856 .expect_err("parent components in workspace should fail closed");
1857
1858 assert!(
1859 format!("{err:#}").contains("workspace path cannot contain '..'"),
1860 "unexpected error: {err:#}"
1861 );
1862 }
1863
1864 #[test]
1865 fn workspace_mcp_config_resolves_relative_cwd_from_workspace() {
1866 let dir = tempfile::tempdir().unwrap();
1867 let global_path = dir.path().join("global-mcp.json");
1868 let workspace = dir.path().join("workspace");
1869 let project_dir = workspace.join(".codewhale");
1870 fs::create_dir_all(&project_dir).unwrap();
1871 let _trust = mark_workspace_trusted(&workspace);
1872 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
1873 fs::write(
1874 project_dir.join("mcp.json"),
1875 r#"{"servers": {"project": {"command": "node", "args": ["server.js"], "cwd": "tools/mcp"}}}"#,
1876 )
1877 .unwrap();
1878
1879 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
1880 let workspace = workspace.canonicalize().unwrap();
1881
1882 let project = cfg.servers.get("project").unwrap();
1883 assert_eq!(
1884 project.cwd.as_deref(),
1885 Some(workspace.join("tools/mcp").as_path())
1886 );
1887 }
1888
1889 #[test]
1890 fn workspace_mcp_config_rejects_project_cwd_escape() {
1891 let dir = tempfile::tempdir().unwrap();
1892 let global_path = dir.path().join("global-mcp.json");
1893 let workspace = dir.path().join("workspace");
1894 let project_dir = workspace.join(".codewhale");
1895 fs::create_dir_all(&project_dir).unwrap();
1896 let _trust = mark_workspace_trusted(&workspace);
1897 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
1898 fs::write(
1899 project_dir.join("mcp.json"),
1900 r#"{"servers": {"project": {"command": "node", "args": ["server.js"], "cwd": "../outside"}}}"#,
1901 )
1902 .unwrap();
1903
1904 let err = load_config_with_workspace(&global_path, &workspace)
1905 .expect_err("project MCP cwd escape must be rejected");
1906
1907 assert!(
1908 err.to_string()
1909 .contains("Project MCP server cwd must stay within workspace"),
1910 "unexpected error: {err}"
1911 );
1912 }
1913
1914 #[cfg(unix)]
1915 #[test]
1916 fn workspace_mcp_config_rejects_symlinked_project_cwd_escape() {
1917 let dir = tempfile::tempdir().unwrap();
1918 let global_path = dir.path().join("global-mcp.json");
1919 let workspace = dir.path().join("workspace");
1920 let project_dir = workspace.join(".codewhale");
1921 let outside = dir.path().join("outside");
1922 fs::create_dir_all(&project_dir).unwrap();
1923 fs::create_dir_all(&outside).unwrap();
1924 std::os::unix::fs::symlink(&outside, workspace.join("tools")).unwrap();
1925 let _trust = mark_workspace_trusted(&workspace);
1926 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
1927 fs::write(
1928 project_dir.join("mcp.json"),
1929 r#"{"servers": {"project": {"command": "node", "args": ["server.js"], "cwd": "tools"}}}"#,
1930 )
1931 .unwrap();
1932
1933 let err = load_config_with_workspace(&global_path, &workspace)
1934 .expect_err("project MCP symlink cwd escape must be rejected");
1935
1936 assert!(
1937 err.to_string()
1938 .contains("Project MCP server cwd must stay within workspace"),
1939 "unexpected error: {err}"
1940 );
1941 }
1942
1943 #[test]
1944 fn workspace_mcp_config_rejects_workspace_traversal() {
1945 let dir = tempfile::tempdir().unwrap();
1946 let global_path = dir.path().join("global-mcp.json");
1947 let workspace = dir.path().join("workspace");
1948 let bad_workspace = workspace.join("..").join("outside");
1949 fs::create_dir_all(&workspace).unwrap();
1950 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
1951
1952 let err = load_config_with_workspace(&global_path, &bad_workspace)
1953 .expect_err("workspace traversal should fail");
1954 assert!(
1955 format!("{err:#}").contains("workspace path cannot contain '..'"),
1956 "unexpected error: {err:#}"
1957 );
1958 }
1959
1960 #[tokio::test]
1961 async fn workspace_mcp_pool_reload_picks_up_project_config_creation() {
1962 let dir = tempfile::tempdir().unwrap();
1963 let global_path = dir.path().join("global-mcp.json");
1964 let workspace = dir.path().join("workspace");
1965 let project_dir = workspace.join(".codewhale");
1966 fs::create_dir_all(&workspace).unwrap();
1967 let _trust = mark_workspace_trusted(&workspace);
1968 fs::write(
1969 &global_path,
1970 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
1971 )
1972 .unwrap();
1973
1974 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
1975 assert_eq!(pool.server_names(), vec!["global".to_string()]);
1976
1977 fs::create_dir_all(&project_dir).unwrap();
1978 fs::write(
1979 project_dir.join("mcp.json"),
1980 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
1981 )
1982 .unwrap();
1983
1984 assert!(pool.reload_if_config_changed().await.unwrap());
1985 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
1986 let expected: std::collections::BTreeSet<String> =
1987 ["global".to_string(), "project".to_string()]
1988 .into_iter()
1989 .collect();
1990 assert_eq!(names, expected);
1991 }
1992
1993 #[tokio::test]
1994 async fn workspace_mcp_pool_reload_picks_up_project_config_after_workspace_trust() {
1995 let dir = tempfile::tempdir().unwrap();
1996 let global_path = dir.path().join("global-mcp.json");
1997 let workspace = dir.path().join("workspace");
1998 let project_dir = workspace.join(".codewhale");
1999 fs::create_dir_all(&project_dir).unwrap();
2000 let trust_env = workspace_trust_config_guard(&workspace);
2001 fs::write(
2002 &global_path,
2003 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2004 )
2005 .unwrap();
2006 fs::write(
2007 project_dir.join("mcp.json"),
2008 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2009 )
2010 .unwrap();
2011
2012 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
2013 assert_eq!(pool.server_names(), vec!["global".to_string()]);
2014
2015 write_workspace_trust_config(&trust_env.config_path, &workspace);
2016
2017 assert!(pool.reload_if_config_changed().await.unwrap());
2018 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
2019 let expected: std::collections::BTreeSet<String> =
2020 ["global".to_string(), "project".to_string()]
2021 .into_iter()
2022 .collect();
2023 assert_eq!(names, expected);
2024 }
2025
2026 #[tokio::test]
2027 async fn workspace_mcp_pool_reload_drops_project_config_after_workspace_trust_removed() {
2028 let dir = tempfile::tempdir().unwrap();
2029 let global_path = dir.path().join("global-mcp.json");
2030 let workspace = dir.path().join("workspace");
2031 let project_dir = workspace.join(".codewhale");
2032 fs::create_dir_all(&project_dir).unwrap();
2033 let trust = mark_workspace_trusted(&workspace);
2034 fs::write(
2035 &global_path,
2036 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2037 )
2038 .unwrap();
2039 fs::write(
2040 project_dir.join("mcp.json"),
2041 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2042 )
2043 .unwrap();
2044
2045 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
2046 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
2047 let expected: std::collections::BTreeSet<String> =
2048 ["global".to_string(), "project".to_string()]
2049 .into_iter()
2050 .collect();
2051 assert_eq!(names, expected);
2052
2053 fs::remove_file(&trust.config_path).unwrap();
2054
2055 assert!(pool.reload_if_config_changed().await.unwrap());
2056 assert_eq!(pool.server_names(), vec!["global".to_string()]);
2057 }
2058
2059 #[tokio::test]
2060 async fn workspace_mcp_pool_reload_drops_project_config_after_deletion() {
2061 let dir = tempfile::tempdir().unwrap();
2062 let global_path = dir.path().join("global-mcp.json");
2063 let workspace = dir.path().join("workspace");
2064 let project_dir = workspace.join(".codewhale");
2065 fs::create_dir_all(&project_dir).unwrap();
2066 let _trust = mark_workspace_trusted(&workspace);
2067 fs::write(
2068 &global_path,
2069 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2070 )
2071 .unwrap();
2072 let project_path = project_dir.join("mcp.json");
2073 fs::write(
2074 &project_path,
2075 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2076 )
2077 .unwrap();
2078
2079 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
2080 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
2081 let expected: std::collections::BTreeSet<String> =
2082 ["global".to_string(), "project".to_string()]
2083 .into_iter()
2084 .collect();
2085 assert_eq!(names, expected);
2086
2087 fs::remove_file(project_path).unwrap();
2088
2089 assert!(pool.reload_if_config_changed().await.unwrap());
2090 assert_eq!(pool.server_names(), vec!["global".to_string()]);
2091 }
2092
2093 #[test]
2094 fn test_mcp_config_rejects_traversal_path() {
2095 let err = load_config(Path::new("../mcp.json")).expect_err("traversal path should fail");
2096 assert!(
2097 format!("{err:#}").contains("cannot contain '..'"),
2098 "got: {err:#}"
2099 );
2100 }
2101
2102 #[cfg(unix)]
2103 #[test]
2104 fn mcp_config_rejects_symlinked_config_file() {
2105 let dir = tempfile::tempdir().unwrap();
2106 let target = dir.path().join("target-mcp.json");
2107 let link = dir.path().join("mcp.json");
2108 fs::write(&target, r#"{"servers": {}}"#).expect("write target config");
2109 std::os::unix::fs::symlink(&target, &link).expect("symlink mcp config");
2110
2111 let err = load_config(&link).expect_err("symlinked MCP config should fail");
2112
2113 assert!(format!("{err:#}").contains("regular file"), "got: {err:#}");
2114 }
2115
2116 #[test]
2117 fn init_mcp_config_rejects_traversal_before_parent_creation() {
2118 let dir = tempfile::tempdir().unwrap();
2119 let outside_dir = dir.path().join("outside");
2120 let path = dir
2121 .path()
2122 .join("allowed")
2123 .join("..")
2124 .join("outside")
2125 .join("mcp.json");
2126
2127 let err = init_config(&path, false).expect_err("traversal path should fail");
2128
2129 assert!(
2130 format!("{err:#}").contains("cannot contain '..'"),
2131 "got: {err:#}"
2132 );
2133 assert!(
2134 !outside_dir.exists(),
2135 "init_config must validate before creating parent directories"
2136 );
2137 }
2138
2139 #[test]
2140 fn test_mcp_config_manager_actions_round_trip() {
2141 let dir = tempfile::tempdir().unwrap();
2142 let path = dir.path().join("mcp.json");
2143
2144 assert_eq!(init_config(&path, false).unwrap(), McpWriteStatus::Created);
2145 assert_eq!(
2146 init_config(&path, false).unwrap(),
2147 McpWriteStatus::SkippedExists
2148 );
2149
2150 add_server_config(
2151 &path,
2152 "local".to_string(),
2153 Some("node".to_string()),
2154 None,
2155 vec!["server.js".to_string()],
2156 None,
2157 )
2158 .unwrap();
2159 set_server_enabled(&path, "local", false).unwrap();
2160 let disabled = manager_snapshot_from_config(&path, true).unwrap();
2161 let local = disabled
2162 .servers
2163 .iter()
2164 .find(|server| server.name == "local")
2165 .unwrap();
2166 assert!(!local.enabled);
2167 assert_eq!(local.transport, "stdio");
2168
2169 remove_server_config(&path, "local").unwrap();
2170 let removed = manager_snapshot_from_config(&path, true).unwrap();
2171 assert!(removed.servers.iter().all(|server| server.name != "local"));
2172 }
2173
2174 #[test]
2175 fn test_mcp_config_adds_explicit_sse_transport() {
2176 let dir = tempfile::tempdir().unwrap();
2177 let path = dir.path().join("mcp.json");
2178
2179 add_server_config(
2180 &path,
2181 "legacy".to_string(),
2182 None,
2183 Some("https://example.com/v1/mcp/sse".to_string()),
2184 Vec::new(),
2185 Some("sse".to_string()),
2186 )
2187 .unwrap();
2188
2189 let cfg = load_config(&path).unwrap();
2190 assert_eq!(
2191 cfg.servers
2192 .get("legacy")
2193 .and_then(|server| server.transport.as_deref()),
2194 Some("sse")
2195 );
2196
2197 let snapshot = manager_snapshot_from_config(&path, false).unwrap();
2198 assert_eq!(snapshot.servers[0].transport, "sse");
2199 }
2200
2201 #[test]
2202 fn test_mcp_config_rejects_unknown_transport() {
2203 let dir = tempfile::tempdir().unwrap();
2204 let path = dir.path().join("mcp.json");
2205
2206 let err = add_server_config(
2207 &path,
2208 "bad".to_string(),
2209 None,
2210 Some("https://example.com/mcp".to_string()),
2211 Vec::new(),
2212 Some("streamable".to_string()),
2213 )
2214 .expect_err("unknown transport should fail");
2215
2216 assert!(
2217 format!("{err:#}").contains("Unsupported MCP transport"),
2218 "got: {err:#}"
2219 );
2220 }
2221
2222 #[test]
2223 fn test_server_effective_timeouts() {
2224 let global = McpTimeouts::default();
2225
2226 let server_with_override = McpServerConfig {
2227 command: Some("test".to_string()),
2228 args: vec![],
2229 env: HashMap::new(),
2230 cwd: None,
2231 url: None,
2232 transport: None,
2233 connect_timeout: Some(20),
2234 execute_timeout: None,
2235 read_timeout: Some(180),
2236 disabled: false,
2237 enabled: true,
2238 required: false,
2239 enabled_tools: Vec::new(),
2240 disabled_tools: Vec::new(),
2241 headers: HashMap::new(),
2242 env_headers: HashMap::new(),
2243 bearer_token_env_var: None,
2244 scopes: Vec::new(),
2245 oauth: None,
2246 oauth_resource: None,
2247 reviewed_plugin: None,
2248 };
2249
2250 assert_eq!(server_with_override.effective_connect_timeout(&global), 20);
2251 assert_eq!(server_with_override.effective_execute_timeout(&global), 60); // global default
2252 assert_eq!(server_with_override.effective_read_timeout(&global), 180);
2253 }
2254
2255 #[test]
2256 fn test_mcp_pool_is_mcp_tool() {
2257 assert!(McpPool::is_mcp_tool("mcp_filesystem_read"));
2258 assert!(McpPool::is_mcp_tool("mcp_git_status"));
2259 assert!(McpPool::is_mcp_tool("list_mcp_resources"));
2260 assert!(McpPool::is_mcp_tool("list_mcp_resource_templates"));
2261 assert!(McpPool::is_mcp_tool("read_mcp_resource"));
2262 assert!(!McpPool::is_mcp_tool("read_file"));
2263 assert!(!McpPool::is_mcp_tool("exec_shell"));
2264 }
2265
2266 struct ScriptedValueTransport {
2267 sent: Arc<Mutex<Vec<serde_json::Value>>>,
2268 responses: VecDeque<Vec<u8>>,
2269 }
2270
2271 #[async_trait::async_trait]
2272 impl McpTransport for ScriptedValueTransport {
2273 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
2274 self.sent
2275 .lock()
2276 .unwrap()
2277 .push(serde_json::from_slice(&msg)?);
2278 Ok(())
2279 }
2280
2281 async fn recv(&mut self) -> Result<Vec<u8>> {
2282 self.responses
2283 .pop_front()
2284 .context("scripted transport exhausted")
2285 }
2286 }
2287
2288 struct HangingValueTransport {
2289 sent: Arc<Mutex<Vec<serde_json::Value>>>,
2290 }
2291
2292 #[async_trait::async_trait]
2293 impl McpTransport for HangingValueTransport {
2294 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
2295 self.sent
2296 .lock()
2297 .unwrap()
2298 .push(serde_json::from_slice(&msg)?);
2299 Ok(())
2300 }
2301
2302 async fn recv(&mut self) -> Result<Vec<u8>> {
2303 std::future::pending().await
2304 }
2305 }
2306
2307 struct ScriptedThenHangingTransport {
2308 sent: Arc<Mutex<Vec<serde_json::Value>>>,
2309 responses: VecDeque<Vec<u8>>,
2310 }
2311
2312 #[async_trait::async_trait]
2313 impl McpTransport for ScriptedThenHangingTransport {
2314 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
2315 self.sent
2316 .lock()
2317 .unwrap()
2318 .push(serde_json::from_slice(&msg)?);
2319 Ok(())
2320 }
2321
2322 async fn recv(&mut self) -> Result<Vec<u8>> {
2323 match self.responses.pop_front() {
2324 Some(response) => Ok(response),
2325 None => std::future::pending().await,
2326 }
2327 }
2328 }
2329
2330 struct DropCountingTransport {
2331 drops: Arc<AtomicUsize>,
2332 }
2333
2334 #[async_trait::async_trait]
2335 impl McpTransport for DropCountingTransport {
2336 async fn send(&mut self, _msg: Vec<u8>) -> Result<()> {
2337 Ok(())
2338 }
2339
2340 async fn recv(&mut self) -> Result<Vec<u8>> {
2341 std::future::pending().await
2342 }
2343 }
2344
2345 impl Drop for DropCountingTransport {
2346 fn drop(&mut self) {
2347 self.drops.fetch_add(1, AtomicOrdering::SeqCst);
2348 }
2349 }
2350
2351 fn test_server_config() -> McpServerConfig {
2352 McpServerConfig {
2353 command: Some("mock".to_string()),
2354 args: Vec::new(),
2355 env: HashMap::new(),
2356 cwd: None,
2357 url: None,
2358 transport: None,
2359 connect_timeout: None,
2360 execute_timeout: None,
2361 read_timeout: None,
2362 disabled: false,
2363 enabled: true,
2364 required: false,
2365 enabled_tools: Vec::new(),
2366 disabled_tools: Vec::new(),
2367 headers: HashMap::new(),
2368 env_headers: HashMap::new(),
2369 bearer_token_env_var: None,
2370 scopes: Vec::new(),
2371 oauth: None,
2372 oauth_resource: None,
2373 reviewed_plugin: None,
2374 }
2375 }
2376
2377 fn test_connection(transport: Box<dyn McpTransport>) -> McpConnection {
2378 McpConnection {
2379 name: "mock".to_string(),
2380 transport,
2381 tools: Vec::new(),
2382 resources: Vec::new(),
2383 resource_templates: Vec::new(),
2384 prompts: Vec::new(),
2385 request_id: AtomicU64::new(1),
2386 state: ConnectionState::Ready,
2387 config: test_server_config(),
2388 server_capabilities: None,
2389 discovery_timeout: Duration::from_secs(default_connect_timeout()),
2390 read_timeout_secs: default_read_timeout(),
2391 cancel_token: tokio_util::sync::CancellationToken::new(),
2392 authority_revocation_reason: Arc::new(std::sync::Mutex::new(None)),
2393 authority_watch: None,
2394 catalog_generation: 0,
2395 }
2396 }
2397
2398 fn json_frame(value: serde_json::Value) -> Vec<u8> {
2399 serde_json::to_vec(&value).unwrap()
2400 }
2401
2402 #[tokio::test]
2403 async fn call_method_skips_notifications_and_unmatched_responses() {
2404 let sent = Arc::new(Mutex::new(Vec::new()));
2405 let transport = ScriptedValueTransport {
2406 sent: Arc::clone(&sent),
2407 responses: VecDeque::from([
2408 json_frame(serde_json::json!({
2409 "jsonrpc": "2.0",
2410 "method": "notifications/progress",
2411 "params": {"progress": 0.5}
2412 })),
2413 json_frame(serde_json::json!({
2414 "jsonrpc": "2.0",
2415 "id": 99,
2416 "result": {"ignored": true}
2417 })),
2418 json_frame(serde_json::json!({
2419 "jsonrpc": "2.0",
2420 "id": 1,
2421 "result": {"ok": true}
2422 })),
2423 ]),
2424 };
2425 let mut conn = test_connection(Box::new(transport));
2426
2427 let result = conn
2428 .call_method("tools/call", serde_json::json!({"name": "echo"}), 1)
2429 .await
2430 .unwrap();
2431
2432 assert_eq!(result, serde_json::json!({"ok": true}));
2433 let sent = sent.lock().unwrap();
2434 assert_eq!(sent.len(), 1);
2435 assert_eq!(sent[0]["jsonrpc"], "2.0");
2436 assert_eq!(sent[0]["id"], "1");
2437 assert_eq!(sent[0]["method"], "tools/call");
2438 }
2439
2440 #[tokio::test]
2441 async fn call_method_invalid_json_includes_server_output_preview() {
2442 let sent = Arc::new(Mutex::new(Vec::new()));
2443 let transport = ScriptedValueTransport {
2444 sent: Arc::clone(&sent),
2445 responses: VecDeque::from([b"Allow Burp MCP connection? [y/N]".to_vec()]),
2446 };
2447 let mut conn = test_connection(Box::new(transport));
2448
2449 let err = conn
2450 .call_method("tools/call", serde_json::json!({"name": "burp"}), 1)
2451 .await
2452 .expect_err("non-json MCP stdout should fail");
2453 let msg = err.to_string();
2454
2455 assert!(msg.contains("Invalid MCP JSON-RPC message from server 'mock'"));
2456 assert!(msg.contains("Allow Burp MCP connection"));
2457 assert_eq!(conn.state(), ConnectionState::Disconnected);
2458 }
2459
2460 #[tokio::test]
2461 async fn recv_times_out_waiting_for_mcp_response_and_disconnects() {
2462 let sent = Arc::new(Mutex::new(Vec::new()));
2463 let mut conn = test_connection(Box::new(HangingValueTransport {
2464 sent: Arc::clone(&sent),
2465 }));
2466 conn.read_timeout_secs = 0;
2467
2468 let err = conn
2469 .recv("1".to_string())
2470 .await
2471 .expect_err("hung transport should time out inside recv");
2472
2473 assert!(
2474 err.to_string()
2475 .contains("Timed out waiting for MCP JSON-RPC response from server 'mock' after 0s"),
2476 "unexpected error: {err:#}"
2477 );
2478 assert_eq!(conn.state(), ConnectionState::Disconnected);
2479 }
2480
2481 #[tokio::test]
2482 async fn call_method_times_out_while_waiting_for_response() {
2483 let sent = Arc::new(Mutex::new(Vec::new()));
2484 let mut conn = test_connection(Box::new(HangingValueTransport {
2485 sent: Arc::clone(&sent),
2486 }));
2487
2488 let err = conn
2489 .call_method("tools/call", serde_json::json!({"name": "echo"}), 0)
2490 .await
2491 .expect_err("hung receive should time out");
2492
2493 assert!(
2494 err.to_string()
2495 .contains("MCP method 'tools/call' on server 'mock' timed out after 0s"),
2496 "unexpected error: {err:#}"
2497 );
2498 assert_eq!(sent.lock().unwrap().len(), 1);
2499 }
2500
2501 #[tokio::test]
2502 async fn test_mcp_pool_empty_config() {
2503 let pool = McpPool::new(McpConfig::default());
2504 assert!(pool.server_names().is_empty());
2505 assert!(pool.all_tools().is_empty());
2506 }
2507
2508 /// #1267 part 2: a pool built without a source path has no file to watch,
2509 /// so `reload_if_config_changed` must short-circuit instead of trying
2510 /// to stat `/`.
2511 #[tokio::test]
2512 async fn reload_if_config_changed_is_noop_without_source_path() {
2513 let mut pool = McpPool::new(McpConfig::default());
2514 let reloaded = pool.reload_if_config_changed().await.unwrap();
2515 assert!(!reloaded, "no source path → no reload");
2516 }
2517
2518 /// #1267 part 2: when the on-disk config is byte-unchanged, the lazy
2519 /// reload must not drop connections — every call to `get_or_connect`
2520 /// would otherwise pay a full reconnect cycle on networked filesystems
2521 /// where mtime granularity is coarse.
2522 #[tokio::test]
2523 async fn reload_if_config_changed_skips_when_content_unchanged() {
2524 let dir = tempfile::tempdir().unwrap();
2525 let path = dir.path().join("mcp.json");
2526 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
2527 let mut pool = McpPool::from_config_path(&path).unwrap();
2528 // Force the mtime to advance without changing content.
2529 std::thread::sleep(std::time::Duration::from_millis(10));
2530 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
2531 let reloaded = pool.reload_if_config_changed().await.unwrap();
2532 assert!(
2533 !reloaded,
2534 "content-unchanged config must not trigger a reload"
2535 );
2536 }
2537
2538 /// #1267 part 2: when the on-disk config changes content, the next
2539 /// `reload_if_config_changed` call must swap in the new config and
2540 /// (would) drop all live connections. We can't stand up a real
2541 /// `McpConnection` in a unit test, so we observe the swap via the
2542 /// publicly-readable side: server names go from empty to non-empty.
2543 #[tokio::test]
2544 async fn reload_if_config_changed_swaps_config_on_content_change() {
2545 let dir = tempfile::tempdir().unwrap();
2546 let path = dir.path().join("mcp.json");
2547 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
2548 let mut pool = McpPool::from_config_path(&path).unwrap();
2549 assert!(pool.server_names().is_empty());
2550 // Mutate the file so both the mtime and the hash change.
2551 std::thread::sleep(std::time::Duration::from_millis(10));
2552 std::fs::write(
2553 &path,
2554 r#"{"servers":{"new":{"command":"echo","args":["hi"]}}}"#,
2555 )
2556 .unwrap();
2557 let reloaded = pool.reload_if_config_changed().await.unwrap();
2558 assert!(reloaded, "content-changed config must trigger reload");
2559 let names = pool.server_names();
2560 assert!(
2561 names.contains(&"new".to_string()),
2562 "expected new server in pool after reload, got {names:?}"
2563 );
2564 }
2565
2566 #[tokio::test]
2567 async fn reload_if_config_changed_drops_live_connections() {
2568 let dir = tempfile::tempdir().unwrap();
2569 let path = dir.path().join("mcp.json");
2570 std::fs::write(
2571 &path,
2572 r#"{"servers":{"local":{"command":"node","args":["server.js"]}}}"#,
2573 )
2574 .unwrap();
2575 let mut pool = McpPool::from_config_path(&path).unwrap();
2576 let drops = Arc::new(AtomicUsize::new(0));
2577 let mut conn = test_connection(Box::new(DropCountingTransport {
2578 drops: Arc::clone(&drops),
2579 }));
2580 conn.name = "local".to_string();
2581 conn.config = pool.config.servers.get("local").unwrap().clone();
2582 pool.connections.insert("local".to_string(), conn);
2583
2584 std::thread::sleep(std::time::Duration::from_millis(10));
2585 std::fs::write(
2586 &path,
2587 r#"{"servers":{"local":{"command":"node","args":["server-v2.js"]}}}"#,
2588 )
2589 .unwrap();
2590
2591 let reloaded = pool.reload_if_config_changed().await.unwrap();
2592 assert!(reloaded, "content-changed config must trigger reload");
2593 assert_eq!(
2594 drops.load(AtomicOrdering::SeqCst),
2595 1,
2596 "reload must drop the stale live transport"
2597 );
2598 assert!(
2599 !pool.connections.contains_key("local"),
2600 "stale connection must not survive config reload"
2601 );
2602 assert_eq!(
2603 pool.config.servers.get("local").unwrap().args,
2604 vec!["server-v2.js".to_string()]
2605 );
2606 }
2607
2608 #[tokio::test]
2609 async fn connect_all_reloads_before_snapshotting_new_server_names() {
2610 let dir = tempfile::tempdir().unwrap();
2611 let path = dir.path().join("mcp.json");
2612 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
2613 let mut pool = McpPool::from_config_path(&path).unwrap();
2614
2615 std::fs::write(
2616 &path,
2617 r#"{"servers":{"late":{"command":"codewhale-test-command-that-does-not-exist"}}}"#,
2618 )
2619 .unwrap();
2620 // Make the test independent of filesystem mtime granularity.
2621 pool.last_mtimes = vec![None];
2622
2623 let errors = pool.connect_all().await;
2624 assert!(
2625 pool.server_names().contains(&"late".to_string()),
2626 "the first connect_all call must install the changed config"
2627 );
2628 assert!(
2629 errors.iter().any(|(name, _)| name == "late"),
2630 "the newly-added server must be attempted on the same call"
2631 );
2632 }
2633
2634 #[tokio::test]
2635 async fn explicit_reload_reconnects_unchanged_config_and_preserves_dynamic_servers() {
2636 let dir = tempfile::tempdir().unwrap();
2637 let path = dir.path().join("mcp.json");
2638 std::fs::write(
2639 &path,
2640 r#"{"servers":{"local":{"command":"node","disabled":true}}}"#,
2641 )
2642 .unwrap();
2643 let mut pool = McpPool::from_config_path(&path).unwrap();
2644 let drops = Arc::new(AtomicUsize::new(0));
2645 let mut conn = test_connection(Box::new(DropCountingTransport {
2646 drops: Arc::clone(&drops),
2647 }));
2648 conn.name = "local".to_string();
2649 conn.config = pool.config.servers.get("local").unwrap().clone();
2650 pool.connections.insert("local".to_string(), conn);
2651 let mut runtime_config = test_server_config();
2652 runtime_config.command = Some("runtime-server".to_string());
2653 pool.add_runtime_server_config("runtime".to_string(), runtime_config)
2654 .unwrap();
2655 let generation_before = pool.catalog_generation.load(AtomicOrdering::SeqCst);
2656
2657 let errors = pool.reload_and_connect_all().await.unwrap();
2658
2659 assert!(
2660 errors.is_empty(),
2661 "disabled config should not connect: {errors:?}"
2662 );
2663 assert_eq!(drops.load(AtomicOrdering::SeqCst), 1);
2664 assert!(!pool.connections.contains_key("local"));
2665 assert!(pool.server_names().contains(&"runtime".to_string()));
2666 assert_eq!(
2667 pool.catalog_generation.load(AtomicOrdering::SeqCst),
2668 generation_before + 1,
2669 "explicit reload must invalidate every previously advertised route"
2670 );
2671 }
2672
2673 #[tokio::test]
2674 async fn config_source_switch_preserves_dynamic_servers_in_the_shared_pool() {
2675 let dir = tempfile::tempdir().unwrap();
2676 let workspace = dir.path().join("workspace");
2677 std::fs::create_dir_all(&workspace).unwrap();
2678 let initial_path = dir.path().join("initial.json");
2679 let invalid_path = dir.path().join("invalid.json");
2680 let replacement_path = dir.path().join("replacement.json");
2681 std::fs::write(
2682 &initial_path,
2683 r#"{"servers":{"local":{"command":"node","disabled":true}}}"#,
2684 )
2685 .unwrap();
2686 std::fs::write(&invalid_path, r#"{"servers":{"broken": trailing}}"#).unwrap();
2687 std::fs::write(&replacement_path, r#"{"servers":{}}"#).unwrap();
2688 let plugins = Arc::new(crate::plugins::PluginRegistry::empty(&workspace));
2689 let mut pool = McpPool::from_config_path_with_workspace_and_plugins(
2690 &initial_path,
2691 &workspace,
2692 Arc::clone(&plugins),
2693 )
2694 .unwrap();
2695 let mut runtime_config = test_server_config();
2696 runtime_config.command = Some("runtime-server".to_string());
2697 pool.add_runtime_server_config("runtime".to_string(), runtime_config)
2698 .unwrap();
2699 let drops = Arc::new(AtomicUsize::new(0));
2700 let mut conn = test_connection(Box::new(DropCountingTransport {
2701 drops: Arc::clone(&drops),
2702 }));
2703 conn.name = "local".to_string();
2704 conn.config = pool.config.servers.get("local").unwrap().clone();
2705 pool.connections.insert("local".to_string(), conn);
2706 let generation_before = pool.catalog_generation.load(AtomicOrdering::SeqCst);
2707
2708 pool.switch_workspace_config_source_and_connect_all(
2709 &invalid_path,
2710 &workspace,
2711 Arc::clone(&plugins),
2712 )
2713 .await
2714 .expect_err("malformed replacement must fail closed");
2715 assert_eq!(pool.config_sources.first(), Some(&initial_path));
2716 assert!(pool.connections.contains_key("local"));
2717 assert_eq!(drops.load(AtomicOrdering::SeqCst), 0);
2718 assert_eq!(
2719 pool.catalog_generation.load(AtomicOrdering::SeqCst),
2720 generation_before
2721 );
2722
2723 let errors = pool
2724 .switch_workspace_config_source_and_connect_all(&replacement_path, &workspace, plugins)
2725 .await
2726 .unwrap();
2727
2728 assert!(errors.is_empty());
2729 assert!(pool.server_names().contains(&"runtime".to_string()));
2730 assert_eq!(pool.config_sources.first(), Some(&replacement_path));
2731 assert_eq!(drops.load(AtomicOrdering::SeqCst), 1);
2732 }
2733
2734 /// #1267 part 2: hash-based comparison must be stable for byte-identical
2735 /// configs and distinct for differing configs.
2736 #[test]
2737 fn hash_mcp_config_is_stable_and_change_sensitive() {
2738 let a = McpConfig::default();
2739 let b = McpConfig::default();
2740 assert_eq!(hash_mcp_config(&a), hash_mcp_config(&b));
2741 let mut c = McpConfig::default();
2742 c.servers.insert(
2743 "x".into(),
2744 McpServerConfig {
2745 command: Some("/bin/echo".into()),
2746 args: vec!["hi".into()],
2747 env: Default::default(),
2748 cwd: None,
2749 url: None,
2750 transport: None,
2751 connect_timeout: None,
2752 execute_timeout: None,
2753 read_timeout: None,
2754 disabled: false,
2755 enabled: true,
2756 required: false,
2757 enabled_tools: Vec::new(),
2758 disabled_tools: Vec::new(),
2759 headers: HashMap::new(),
2760 env_headers: HashMap::new(),
2761 bearer_token_env_var: None,
2762 scopes: Vec::new(),
2763 oauth: None,
2764 oauth_resource: None,
2765 reviewed_plugin: None,
2766 },
2767 );
2768 assert_ne!(
2769 hash_mcp_config(&a),
2770 hash_mcp_config(&c),
2771 "hash must change when servers map changes"
2772 );
2773 }
2774
2775 /// #1319: discovered tools must be sorted by name so the prompt prefix
2776 /// is stable across runs (cache-hit stability), even when the server
2777 /// returns them in arbitrary or paginated order.
2778 #[tokio::test]
2779 async fn discover_tools_sorts_by_name_for_cache_stability() {
2780 let sent = Arc::new(Mutex::new(Vec::new()));
2781 let transport = ScriptedValueTransport {
2782 sent: Arc::clone(&sent),
2783 responses: VecDeque::from([
2784 json_frame(serde_json::json!({
2785 "jsonrpc": "2.0",
2786 "id": 1,
2787 "result": {
2788 "tools": [
2789 { "name": "zeta", "inputSchema": {} },
2790 { "name": "alpha", "inputSchema": {} }
2791 ],
2792 "nextCursor": "page-2"
2793 }
2794 })),
2795 json_frame(serde_json::json!({
2796 "jsonrpc": "2.0",
2797 "id": 2,
2798 "result": {
2799 "tools": [
2800 { "name": "mu", "inputSchema": {} },
2801 { "name": "beta", "inputSchema": {} }
2802 ]
2803 }
2804 })),
2805 ]),
2806 };
2807 let mut conn = test_connection(Box::new(transport));
2808 conn.discover_tools().await.expect("discover");
2809
2810 let names: Vec<&str> = conn.tools.iter().map(|t| t.name.as_str()).collect();
2811 assert_eq!(
2812 names,
2813 vec!["alpha", "beta", "mu", "zeta"],
2814 "tools must be sorted by name regardless of server order or pagination"
2815 );
2816 }
2817
2818 #[tokio::test]
2819 async fn discover_tools_rejects_a_repeated_pagination_cursor_without_publishing_partials() {
2820 let transport = ScriptedValueTransport {
2821 sent: Arc::new(Mutex::new(Vec::new())),
2822 responses: VecDeque::from([
2823 json_frame(serde_json::json!({
2824 "jsonrpc": "2.0",
2825 "id": 1,
2826 "result": {
2827 "tools": [{ "name": "first", "inputSchema": {} }],
2828 "nextCursor": "same"
2829 }
2830 })),
2831 json_frame(serde_json::json!({
2832 "jsonrpc": "2.0",
2833 "id": 2,
2834 "result": {
2835 "tools": [{ "name": "second", "inputSchema": {} }],
2836 "nextCursor": "same"
2837 }
2838 })),
2839 ]),
2840 };
2841 let mut conn = test_connection(Box::new(transport));
2842
2843 let error = conn
2844 .discover_tools()
2845 .await
2846 .expect_err("repeated cursor must abort discovery");
2847 assert!(error.to_string().contains("repeated pagination cursor"));
2848 assert!(
2849 conn.tools.is_empty(),
2850 "an aborted catalogue must not publish attacker-controlled partial entries"
2851 );
2852 }
2853
2854 #[test]
2855 fn mcp_tool_description_formatter_is_one_line_and_unicode_safe() {
2856 let long_cjk = format!("{}\n这行不应显示", "鲸".repeat(81));
2857 assert_eq!(
2858 format_mcp_tool_description(Some(&long_cjk)),
2859 format!(": {}...", "鲸".repeat(80))
2860 );
2861 assert_eq!(
2862 format_mcp_tool_description(Some("第一行\r\n第二行")),
2863 ": 第一行"
2864 );
2865 assert_eq!(format_mcp_tool_description(Some(" \nignored")), "");
2866 assert_eq!(format_mcp_tool_description(None), "");
2867 }
2868
2869 #[tokio::test]
2870 async fn discover_all_honors_tools_only_server_capabilities() {
2871 let sent = Arc::new(Mutex::new(Vec::new()));
2872 let transport = ScriptedValueTransport {
2873 sent: Arc::clone(&sent),
2874 responses: VecDeque::from([
2875 json_frame(serde_json::json!({
2876 "jsonrpc": "2.0",
2877 "id": 1,
2878 "result": {
2879 "protocolVersion": "2024-11-05",
2880 "serverInfo": {"name": "tools-only", "version": "1.0.0"},
2881 "capabilities": {"tools": {}}
2882 }
2883 })),
2884 json_frame(serde_json::json!({
2885 "jsonrpc": "2.0",
2886 "id": 2,
2887 "result": {
2888 "tools": [{"name": "idea_search", "inputSchema": {}}]
2889 }
2890 })),
2891 ]),
2892 };
2893 let mut conn = test_connection(Box::new(transport));
2894
2895 conn.initialize().await.expect("initialize");
2896 conn.discover_all().await.expect("discover tools");
2897
2898 assert_eq!(conn.tools.len(), 1);
2899 assert!(conn.resources.is_empty());
2900 assert!(conn.resource_templates.is_empty());
2901 assert!(conn.prompts.is_empty());
2902 let methods: Vec<_> = sent
2903 .lock()
2904 .unwrap()
2905 .iter()
2906 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
2907 .map(str::to_string)
2908 .collect();
2909 assert_eq!(
2910 methods,
2911 ["initialize", "notifications/initialized", "tools/list"]
2912 );
2913 }
2914
2915 #[tokio::test]
2916 async fn discover_all_populates_every_advertised_capability() {
2917 let sent = Arc::new(Mutex::new(Vec::new()));
2918 let transport = ScriptedValueTransport {
2919 sent: Arc::clone(&sent),
2920 responses: VecDeque::from([
2921 json_frame(serde_json::json!({
2922 "jsonrpc": "2.0",
2923 "id": 1,
2924 "result": {
2925 "protocolVersion": "2024-11-05",
2926 "serverInfo": {"name": "full", "version": "1.0.0"},
2927 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
2928 }
2929 })),
2930 json_frame(serde_json::json!({
2931 "jsonrpc": "2.0",
2932 "id": 2,
2933 "result": {"tools": [{"name": "search", "inputSchema": {}}]}
2934 })),
2935 json_frame(serde_json::json!({
2936 "jsonrpc": "2.0",
2937 "id": 3,
2938 "result": {"resources": [{"uri": "file:///readme", "name": "readme"}]}
2939 })),
2940 json_frame(serde_json::json!({
2941 "jsonrpc": "2.0",
2942 "id": 4,
2943 "result": {
2944 "resourceTemplates": [{"uriTemplate": "file:///{path}", "name": "file"}]
2945 }
2946 })),
2947 json_frame(serde_json::json!({
2948 "jsonrpc": "2.0",
2949 "id": 5,
2950 "result": {"prompts": [{"name": "review"}]}
2951 })),
2952 ]),
2953 };
2954 let mut conn = test_connection(Box::new(transport));
2955
2956 conn.initialize().await.expect("initialize");
2957 conn.discover_all().await.expect("discover all");
2958
2959 assert_eq!(conn.tools.len(), 1);
2960 assert_eq!(conn.resources.len(), 1);
2961 assert_eq!(conn.resource_templates.len(), 1);
2962 assert_eq!(conn.prompts.len(), 1);
2963 let methods: Vec<_> = sent
2964 .lock()
2965 .unwrap()
2966 .iter()
2967 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
2968 .map(str::to_string)
2969 .collect();
2970 assert_eq!(
2971 methods,
2972 [
2973 "initialize",
2974 "notifications/initialized",
2975 "tools/list",
2976 "resources/list",
2977 "resources/templates/list",
2978 "prompts/list",
2979 ]
2980 );
2981 }
2982
2983 #[tokio::test]
2984 async fn legacy_optional_discovery_hangs_are_bounded_and_fail_soft() {
2985 let sent = Arc::new(Mutex::new(Vec::new()));
2986 let transport = ScriptedThenHangingTransport {
2987 sent: Arc::clone(&sent),
2988 responses: VecDeque::from([json_frame(serde_json::json!({
2989 "jsonrpc": "2.0",
2990 "id": 1,
2991 "result": {"tools": [{"name": "search", "inputSchema": {}}]}
2992 }))]),
2993 };
2994 let mut conn = test_connection(Box::new(transport));
2995 conn.discovery_timeout = Duration::from_millis(60);
2996
2997 let started = tokio::time::Instant::now();
2998 conn.discover_all()
2999 .await
3000 .expect("hung optional methods must not fail discovery");
3001
3002 assert_eq!(conn.tools.len(), 1);
3003 assert!(
3004 started.elapsed() < Duration::from_secs(1),
3005 "optional discovery exceeded its bounded budget: {:?}",
3006 started.elapsed()
3007 );
3008 let methods: Vec<_> = sent
3009 .lock()
3010 .unwrap()
3011 .iter()
3012 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
3013 .map(str::to_string)
3014 .collect();
3015 assert_eq!(
3016 methods,
3017 [
3018 "tools/list",
3019 "resources/list",
3020 "resources/templates/list",
3021 "prompts/list",
3022 ]
3023 );
3024 }
3025
3026 #[tokio::test]
3027 async fn mcp_pool_call_tool_preserves_tool_names_with_dashes() {
3028 let sent = Arc::new(Mutex::new(Vec::new()));
3029 let transport = ScriptedValueTransport {
3030 sent: Arc::clone(&sent),
3031 responses: VecDeque::from([json_frame(serde_json::json!({
3032 "jsonrpc": "2.0",
3033 "id": 1,
3034 "result": {"ok": true}
3035 }))]),
3036 };
3037 let mut conn = test_connection(Box::new(transport));
3038 conn.name = "dephy".to_string();
3039 conn.tools = vec![McpTool {
3040 name: "company--search".to_string(),
3041 description: None,
3042 input_schema: serde_json::json!({}),
3043 }];
3044
3045 let mut pool = McpPool::new(McpConfig {
3046 timeouts: McpTimeouts::default(),
3047 servers: HashMap::new(),
3048 });
3049 pool.connections.insert("dephy".to_string(), conn);
3050
3051 let result = pool
3052 .call_tool(
3053 "mcp_dephy_company--search",
3054 serde_json::json!({"query": "dephy"}),
3055 )
3056 .await
3057 .unwrap();
3058
3059 assert_eq!(result, serde_json::json!({"ok": true}));
3060 let sent = sent.lock().unwrap();
3061 assert_eq!(sent[0]["method"], "tools/call");
3062 assert_eq!(sent[0]["params"]["name"], "company--search");
3063 assert_eq!(
3064 sent[0]["params"]["arguments"],
3065 serde_json::json!({"query": "dephy"})
3066 );
3067 }
3068
3069 #[tokio::test]
3070 async fn mcp_pool_rejects_unadvertised_tool_without_sending_tools_call() {
3071 let sent = Arc::new(Mutex::new(Vec::new()));
3072 let transport = ScriptedValueTransport {
3073 sent: Arc::clone(&sent),
3074 // A malicious server could implement this hidden method, but local
3075 // catalog authorization must prevent the transport from seeing it.
3076 responses: VecDeque::from([json_frame(serde_json::json!({
3077 "jsonrpc": "2.0", "id": 1, "result": {"deleted": true}
3078 }))]),
3079 };
3080 let mut conn = test_connection(Box::new(transport));
3081 conn.name = "spy".to_string();
3082 conn.tools = vec![McpTool {
3083 name: "read".to_string(),
3084 description: None,
3085 input_schema: serde_json::json!({}),
3086 }];
3087 let mut pool = McpPool::new(McpConfig::default());
3088 pool.connections.insert("spy".to_string(), conn);
3089
3090 let error = pool
3091 .call_tool("mcp_spy_delete", serde_json::json!({}))
3092 .await
3093 .expect_err("unadvertised hidden tool must fail locally");
3094 assert!(error.to_string().contains("Unknown MCP tool name"));
3095 assert!(sent.lock().unwrap().is_empty(), "zero tools/call requests");
3096 }
3097
3098 #[tokio::test]
3099 async fn mcp_pool_binds_prompts_and_resources_to_advertised_catalog() {
3100 let sent = Arc::new(Mutex::new(Vec::new()));
3101 let transport = ScriptedValueTransport {
3102 sent: Arc::clone(&sent),
3103 responses: VecDeque::from([json_frame(serde_json::json!({
3104 "jsonrpc": "2.0", "id": 1, "result": {"contents": []}
3105 }))]),
3106 };
3107 let mut conn = test_connection(Box::new(transport));
3108 conn.name = "catalog".to_string();
3109 conn.prompts = vec![McpPrompt {
3110 name: "review".to_string(),
3111 description: None,
3112 arguments: Vec::new(),
3113 }];
3114 conn.resources = vec![McpResource {
3115 uri: "file:///readme".to_string(),
3116 name: "readme".to_string(),
3117 description: None,
3118 mime_type: None,
3119 }];
3120 conn.resource_templates = vec![McpResourceTemplate {
3121 uri_template: "repo://item/{id}".to_string(),
3122 name: "item".to_string(),
3123 description: None,
3124 mime_type: None,
3125 }];
3126 let mut pool = McpPool::new(McpConfig::default());
3127 pool.connections.insert("catalog".to_string(), conn);
3128
3129 pool.get_prompt("catalog", "hidden", serde_json::json!({}))
3130 .await
3131 .expect_err("hidden prompt must fail locally");
3132 pool.read_resource("catalog", "file:///hidden")
3133 .await
3134 .expect_err("hidden literal resource must fail locally");
3135 assert!(sent.lock().unwrap().is_empty());
3136
3137 let result = pool
3138 .read_resource("catalog", "repo://item/42")
3139 .await
3140 .expect("exact advertised template expansion is callable");
3141 assert_eq!(result, serde_json::json!({"contents": []}));
3142 let sent = sent.lock().unwrap();
3143 assert_eq!(sent.len(), 1);
3144 assert_eq!(sent[0]["method"], "resources/read");
3145 assert_eq!(sent[0]["params"]["uri"], "repo://item/42");
3146 }
3147
3148 #[tokio::test]
3149 async fn mcp_pool_call_tool_preserves_server_names_with_underscores() {
3150 let sent = Arc::new(Mutex::new(Vec::new()));
3151 let transport = ScriptedValueTransport {
3152 sent: Arc::clone(&sent),
3153 responses: VecDeque::from([json_frame(serde_json::json!({
3154 "jsonrpc": "2.0",
3155 "id": 1,
3156 "result": {"ok": true}
3157 }))]),
3158 };
3159 let mut conn = test_connection(Box::new(transport));
3160 conn.name = "my_db".to_string();
3161 conn.tools = vec![McpTool {
3162 name: "execute_sql".to_string(),
3163 description: None,
3164 input_schema: serde_json::json!({}),
3165 }];
3166
3167 let mut pool = McpPool::new(McpConfig {
3168 timeouts: McpTimeouts::default(),
3169 servers: HashMap::new(),
3170 });
3171 pool.connections.insert("my_db".to_string(), conn);
3172
3173 let result = pool
3174 .call_tool(
3175 "mcp_my_db_execute_sql",
3176 serde_json::json!({"query": "select 1"}),
3177 )
3178 .await
3179 .unwrap();
3180
3181 assert_eq!(result, serde_json::json!({"ok": true}));
3182 let sent = sent.lock().unwrap();
3183 assert_eq!(sent[0]["method"], "tools/call");
3184 assert_eq!(sent[0]["params"]["name"], "execute_sql");
3185 assert_eq!(
3186 sent[0]["params"]["arguments"],
3187 serde_json::json!({"query": "select 1"})
3188 );
3189 }
3190
3191 #[tokio::test]
3192 async fn mcp_pool_hides_and_rejects_ambiguous_model_tool_names() {
3193 let sent_short = Arc::new(Mutex::new(Vec::new()));
3194 let short_transport = ScriptedValueTransport {
3195 sent: Arc::clone(&sent_short),
3196 responses: VecDeque::from([json_frame(serde_json::json!({
3197 "jsonrpc": "2.0",
3198 "id": 1,
3199 "result": {"short": true}
3200 }))]),
3201 };
3202 let mut short_conn = test_connection(Box::new(short_transport));
3203 short_conn.name = "my".to_string();
3204 short_conn.tools = vec![McpTool {
3205 name: "db_execute_sql".to_string(),
3206 description: None,
3207 input_schema: serde_json::json!({}),
3208 }];
3209
3210 let sent_long = Arc::new(Mutex::new(Vec::new()));
3211 let long_transport = ScriptedValueTransport {
3212 sent: Arc::clone(&sent_long),
3213 responses: VecDeque::from([json_frame(serde_json::json!({
3214 "jsonrpc": "2.0",
3215 "id": 1,
3216 "result": {"long": true}
3217 }))]),
3218 };
3219 let mut long_conn = test_connection(Box::new(long_transport));
3220 long_conn.name = "my_db".to_string();
3221 long_conn.tools = vec![McpTool {
3222 name: "execute_sql".to_string(),
3223 description: None,
3224 input_schema: serde_json::json!({}),
3225 }];
3226
3227 let mut pool = McpPool::new(McpConfig {
3228 timeouts: McpTimeouts::default(),
3229 servers: HashMap::new(),
3230 });
3231 pool.connections.insert("my".to_string(), short_conn);
3232 pool.connections.insert("my_db".to_string(), long_conn);
3233
3234 assert!(
3235 pool.all_tools().is_empty(),
3236 "ambiguous names must never be advertised to the model"
3237 );
3238 let error = pool
3239 .call_tool(
3240 "mcp_my_db_execute_sql",
3241 serde_json::json!({"query": "select 1"}),
3242 )
3243 .await
3244 .expect_err("ambiguous tool route must fail closed");
3245
3246 assert!(error.to_string().contains("Ambiguous MCP tool name"));
3247 assert!(
3248 sent_short.lock().unwrap().is_empty(),
3249 "neither authority may receive an ambiguous tool call"
3250 );
3251 assert!(
3252 sent_long.lock().unwrap().is_empty(),
3253 "neither authority may receive an ambiguous tool call"
3254 );
3255 }
3256
3257 #[tokio::test]
3258 async fn json_rpc_session_error_is_marked_stale() {
3259 let sent = Arc::new(Mutex::new(Vec::new()));
3260 let transport = ScriptedValueTransport {
3261 sent: Arc::clone(&sent),
3262 responses: VecDeque::from([json_frame(serde_json::json!({
3263 "jsonrpc": "2.0",
3264 "id": 1,
3265 "error": {
3266 "code": -32001,
3267 "message": "MCP session expired"
3268 }
3269 }))]),
3270 };
3271 let mut conn = test_connection(Box::new(transport));
3272
3273 let err = conn
3274 .call_tool("search", serde_json::json!({"query": "dephy"}), 1)
3275 .await
3276 .expect_err("session error should fail");
3277
3278 assert!(
3279 is_mcp_stale_session_error(&err),
3280 "JSON-RPC session error should be retryable, got: {err:#}"
3281 );
3282 }
3283
3284 #[test]
3285 fn sse_transport_closed_is_retryable() {
3286 let err = anyhow::anyhow!("SSE transport closed");
3287 assert!(
3288 is_mcp_stale_session_error(&err),
3289 "closed SSE stream should force reconnect before retry"
3290 );
3291 }
3292
3293 #[test]
3294 fn legacy_sse_post_disconnect_is_retryable() {
3295 let err = anyhow::anyhow!(
3296 "MCP SSE POST send failed (transport=sse endpoint=http://127.0.0.1:123/messages): connection closed before message completed"
3297 );
3298 assert!(
3299 is_mcp_stale_session_error(&err),
3300 "closed legacy SSE POST should force reconnect before retry"
3301 );
3302
3303 let err = anyhow::anyhow!(
3304 "MCP SSE POST send failed (transport=sse endpoint=http://127.0.0.1:123/messages): connection reset by peer"
3305 );
3306 assert!(
3307 is_mcp_stale_session_error(&err),
3308 "reset legacy SSE POST should force reconnect before retry"
3309 );
3310
3311 let err = anyhow::anyhow!(
3312 "MCP SSE POST send failed (transport=sse endpoint=http://127.0.0.1:123/messages): An existing connection was forcibly closed by the remote host."
3313 );
3314 assert!(
3315 is_mcp_stale_session_error(&err),
3316 "Windows reset wording should force reconnect before retry"
3317 );
3318 }
3319
3320 #[tokio::test]
3321 async fn discover_all_ignores_unsupported_optional_capabilities() {
3322 let sent = Arc::new(Mutex::new(Vec::new()));
3323 let transport = ScriptedValueTransport {
3324 sent: Arc::clone(&sent),
3325 responses: VecDeque::from([
3326 json_frame(serde_json::json!({
3327 "jsonrpc": "2.0",
3328 "id": 1,
3329 "result": {
3330 "tools": [
3331 { "name": "search", "inputSchema": {} }
3332 ]
3333 }
3334 })),
3335 json_frame(serde_json::json!({
3336 "jsonrpc": "2.0",
3337 "id": 2,
3338 "error": {
3339 "code": -32601,
3340 "message": "resources not supported"
3341 }
3342 })),
3343 json_frame(serde_json::json!({
3344 "jsonrpc": "2.0",
3345 "id": 3,
3346 "error": {
3347 "code": -32601,
3348 "message": "resource templates not supported"
3349 }
3350 })),
3351 json_frame(serde_json::json!({
3352 "jsonrpc": "2.0",
3353 "id": 4,
3354 "error": {
3355 "code": -32601,
3356 "message": "prompts not supported"
3357 }
3358 })),
3359 ]),
3360 };
3361 let mut conn = test_connection(Box::new(transport));
3362 conn.server_capabilities = Some(McpServerCapabilities {
3363 tools: true,
3364 resources: true,
3365 prompts: true,
3366 });
3367
3368 conn.discover_all().await.expect("discover");
3369
3370 assert_eq!(conn.tools.len(), 1);
3371 assert_eq!(conn.tools[0].name, "search");
3372 assert!(conn.resources.is_empty());
3373 assert!(conn.resource_templates.is_empty());
3374 assert!(conn.prompts.is_empty());
3375 let methods: Vec<_> = sent
3376 .lock()
3377 .unwrap()
3378 .iter()
3379 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
3380 .map(str::to_string)
3381 .collect();
3382 assert_eq!(
3383 methods,
3384 [
3385 "tools/list",
3386 "resources/list",
3387 "resources/templates/list",
3388 "prompts/list",
3389 ]
3390 );
3391 }
3392
3393 #[tokio::test]
3394 async fn discover_all_keeps_advertised_tool_discovery_required() {
3395 let sent = Arc::new(Mutex::new(Vec::new()));
3396 let transport = ScriptedValueTransport {
3397 sent,
3398 responses: VecDeque::from([json_frame(serde_json::json!({
3399 "jsonrpc": "2.0",
3400 "id": 1,
3401 "error": {"code": -32601, "message": "tools not supported"}
3402 }))]),
3403 };
3404 let mut conn = test_connection(Box::new(transport));
3405 conn.server_capabilities = Some(McpServerCapabilities {
3406 tools: true,
3407 resources: false,
3408 prompts: false,
3409 });
3410
3411 let error = conn
3412 .discover_all()
3413 .await
3414 .expect_err("advertised tools/list failure must fail discovery");
3415
3416 assert!(
3417 error.to_string().contains("MCP error in 'tools/list'"),
3418 "unexpected error: {error:#}"
3419 );
3420 }
3421
3422 /// #1244: when an MCP stdio server fails to spawn, the underlying OS
3423 /// error (e.g. ENOENT for a missing binary) must reach the user via the
3424 /// snapshot.error string. Regression test for `err.to_string()` dropping
3425 /// the anyhow chain — without `{err:#}` the user sees only the opaque
3426 /// wrapper "MCP stdio spawn failed (...)" and has nothing to act on.
3427 #[tokio::test]
3428 async fn discover_snapshot_includes_underlying_spawn_error_in_chain() {
3429 let dir = tempfile::tempdir().unwrap();
3430 let path = dir.path().join("mcp.json");
3431 fs::write(
3432 &path,
3433 r#"{
3434 "mcpServers": {
3435 "broken": {
3436 "command": "codewhale-tui-test-this-binary-does-not-exist-9f8e7d6c5b4a",
3437 "args": []
3438 }
3439 }
3440 }"#,
3441 )
3442 .unwrap();
3443
3444 let snapshot = discover_manager_snapshot(&path, None, false).await.unwrap();
3445 let server = snapshot
3446 .servers
3447 .iter()
3448 .find(|s| s.name == "broken")
3449 .expect("broken server should appear in snapshot");
3450 let err = server
3451 .error
3452 .as_deref()
3453 .expect("broken server should have an error");
3454 let lowered = err.to_lowercase();
3455 assert!(
3456 lowered.contains("os error")
3457 || lowered.contains("not found")
3458 || lowered.contains("no such"),
3459 "expected underlying spawn error in chain, got: {err}"
3460 );
3461 }
3462
3463 #[test]
3464 fn parse_sse_message_data_extracts_message_events() {
3465 let body = "event: message\r\ndata: {\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\r\n\r\n";
3466 let messages = parse_sse_message_data(body);
3467 assert_eq!(messages.len(), 1);
3468 let value: serde_json::Value = serde_json::from_slice(&messages[0]).unwrap();
3469 assert_eq!(value["id"], 1);
3470 assert!(value.get("result").is_some());
3471 }
3472
3473 #[test]
3474 fn response_id_matches_string_and_numeric_echoes() {
3475 assert!(response_id_matches(Some(&serde_json::json!("1")), "1"));
3476 assert!(response_id_matches(Some(&serde_json::json!(1)), "1"));
3477 assert!(!response_id_matches(Some(&serde_json::json!("2")), "1"));
3478 }
3479
3480 #[test]
3481 fn legacy_sse_transport_requires_explicit_config() {
3482 let mut server = test_server_config();
3483 server.url = Some("https://example.com/mcp/abc/sse".to_string());
3484
3485 assert!(
3486 !is_legacy_sse_transport(&server),
3487 "/sse paths must not force legacy SSE without an explicit transport override"
3488 );
3489
3490 server.transport = Some("sse".to_string());
3491 assert!(is_legacy_sse_transport(&server));
3492
3493 server.transport = Some("SSE".to_string());
3494 assert!(is_legacy_sse_transport(&server));
3495
3496 server.transport = Some("http".to_string());
3497 assert!(!is_legacy_sse_transport(&server));
3498 }
3499
3500 #[test]
3501 fn find_sse_event_separator_accepts_lf_and_crlf() {
3502 assert_eq!(
3503 find_sse_event_separator("event: endpoint\n\n"),
3504 Some((15, 2))
3505 );
3506 assert_eq!(
3507 find_sse_event_separator("event: endpoint\r\n\r\n"),
3508 Some((15, 4))
3509 );
3510 }
3511
3512 #[test]
3513 fn find_sse_event_separator_bytes_matches_str_and_survives_multibyte() {
3514 // Same offsets as the str version.
3515 assert_eq!(
3516 find_sse_event_separator_bytes(b"event: endpoint\n\n"),
3517 Some((15, 2))
3518 );
3519 assert_eq!(
3520 find_sse_event_separator_bytes(b"event: endpoint\r\n\r\n"),
3521 Some((15, 4))
3522 );
3523 // A frame whose data holds a multi-byte char, accumulated byte-wise and
3524 // split mid-char across two reads, decodes intact (no U+FFFD).
3525 let frame = "data: 你好\n\n";
3526 let bytes = frame.as_bytes();
3527 let split = bytes.len() - 3; // inside "好" / before the separator
3528 let mut buffer: Vec<u8> = Vec::new();
3529 buffer.extend_from_slice(&bytes[..split]);
3530 assert_eq!(find_sse_event_separator_bytes(&buffer), None);
3531 buffer.extend_from_slice(&bytes[split..]);
3532 let (pos, sep) = find_sse_event_separator_bytes(&buffer).expect("separator");
3533 let block = String::from_utf8_lossy(&buffer[..pos]).into_owned();
3534 assert_eq!(block, "data: 你好");
3535 assert!(!block.contains('\u{FFFD}'), "multibyte corrupted");
3536 assert_eq!(sep, 2);
3537 }
3538
3539 #[tokio::test]
3540 #[ignore = "flaky: requires a live TCP listener and is sensitive to port allocation races"]
3541 async fn mcp_connection_supports_streamable_http_event_stream_responses() {
3542 use tokio::io::{AsyncReadExt, AsyncWriteExt};
3543 use tokio::net::{TcpListener, TcpStream};
3544
3545 async fn read_http_request(socket: &mut TcpStream) -> String {
3546 let mut request = Vec::new();
3547 let mut buf = [0; 1024];
3548 let header_end = loop {
3549 let n = socket.read(&mut buf).await.unwrap();
3550 assert!(n > 0, "client closed before headers completed");
3551 request.extend_from_slice(&buf[..n]);
3552 if let Some(pos) = request.windows(4).position(|window| window == b"\r\n\r\n") {
3553 break pos + 4;
3554 }
3555 };
3556
3557 let headers = String::from_utf8_lossy(&request[..header_end]);
3558 let content_length = headers
3559 .lines()
3560 .find_map(|line| {
3561 let (name, value) = line.split_once(':')?;
3562 name.eq_ignore_ascii_case("content-length")
3563 .then(|| value.trim().parse::<usize>().ok())
3564 .flatten()
3565 })
3566 .unwrap_or(0);
3567 let total_len = header_end + content_length;
3568 while request.len() < total_len {
3569 let n = socket.read(&mut buf).await.unwrap();
3570 assert!(n > 0, "client closed before body completed");
3571 request.extend_from_slice(&buf[..n]);
3572 }
3573
3574 String::from_utf8(request).unwrap()
3575 }
3576
3577 async fn write_json_sse(socket: &mut TcpStream, response: serde_json::Value) {
3578 let body = format!("event: message\ndata: {response}\n\n");
3579 let response = format!(
3580 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\n\r\n{}",
3581 body.len(),
3582 body
3583 );
3584 socket.write_all(response.as_bytes()).await.unwrap();
3585 }
3586
3587 let _lock = lock_mcp_loopback_tests().await;
3588 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3589 let addr = listener.local_addr().unwrap();
3590 let server = tokio::spawn(async move {
3591 loop {
3592 let Ok((mut socket, _)) = listener.accept().await else {
3593 break;
3594 };
3595 tokio::spawn(async move {
3596 let request = read_http_request(&mut socket).await;
3597 assert!(request.starts_with("POST /mcp "));
3598 assert!(
3599 request.contains("Accept: application/json, text/event-stream")
3600 || request.contains("accept: application/json, text/event-stream")
3601 );
3602 let body = request.split("\r\n\r\n").nth(1).unwrap_or("");
3603 let value: serde_json::Value = serde_json::from_str(body).unwrap();
3604 let method = value["method"].as_str().unwrap();
3605
3606 if method == "notifications/initialized" {
3607 socket
3608 .write_all(b"HTTP/1.1 202 Accepted\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
3609 .await
3610 .unwrap();
3611 return;
3612 }
3613
3614 let id = value["id"].clone();
3615 let result = match method {
3616 "initialize" => serde_json::json!({
3617 "protocolVersion": "2024-11-05",
3618 "serverInfo": {"name": "mock-streamable", "version": "1.0.0"},
3619 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
3620 }),
3621 "tools/list" => serde_json::json!({
3622 "tools": [{
3623 "name": "read_wiki_structure",
3624 "description": "Read wiki structure",
3625 "inputSchema": {"type": "object"}
3626 }]
3627 }),
3628 "resources/list" => serde_json::json!({"resources": []}),
3629 "resources/templates/list" => {
3630 serde_json::json!({"resourceTemplates": []})
3631 }
3632 "prompts/list" => serde_json::json!({"prompts": []}),
3633 other => panic!("unexpected method: {other}"),
3634 };
3635 write_json_sse(
3636 &mut socket,
3637 serde_json::json!({
3638 "jsonrpc": "2.0",
3639 "id": id,
3640 "result": result
3641 }),
3642 )
3643 .await;
3644 });
3645 }
3646 });
3647
3648 let config = McpServerConfig {
3649 command: None,
3650 args: vec![],
3651 env: HashMap::new(),
3652 cwd: None,
3653 url: Some(format!("http://{addr}/mcp")),
3654 transport: None,
3655 connect_timeout: Some(2),
3656 execute_timeout: None,
3657 read_timeout: None,
3658 disabled: false,
3659 enabled: true,
3660 required: false,
3661 enabled_tools: Vec::new(),
3662 disabled_tools: Vec::new(),
3663 headers: HashMap::new(),
3664 env_headers: HashMap::new(),
3665 bearer_token_env_var: None,
3666 scopes: Vec::new(),
3667 oauth: None,
3668 oauth_resource: None,
3669 reviewed_plugin: None,
3670 };
3671
3672 let conn = McpConnection::connect_with_policy(
3673 "deepwiki".to_string(),
3674 config,
3675 &McpTimeouts::default(),
3676 None,
3677 )
3678 .await
3679 .unwrap();
3680
3681 assert_eq!(conn.state(), ConnectionState::Ready);
3682 assert_eq!(conn.tools().len(), 1);
3683 assert_eq!(conn.tools()[0].name, "read_wiki_structure");
3684
3685 server.abort();
3686 }
3687
3688 #[test]
3689 fn mask_url_secrets_strips_userinfo() {
3690 let masked = mask_url_secrets("https://user:s3cret@host.example/api?foo=bar");
3691 assert!(masked.contains("***"), "expected masked userinfo: {masked}");
3692 assert!(!masked.contains("s3cret"), "secret leaked: {masked}");
3693 assert!(masked.contains("host.example"), "host preserved: {masked}");
3694 }
3695
3696 #[test]
3697 fn mask_url_secrets_passes_through_clean_url() {
3698 assert_eq!(
3699 mask_url_secrets("https://api.example.com/mcp"),
3700 "https://api.example.com/mcp"
3701 );
3702 }
3703
3704 #[test]
3705 fn redact_body_preview_masks_bearer_token() {
3706 let redacted = redact_body_preview(
3707 "Authorization: Bearer abc.def.ghi end; authorization: bearer second-token end",
3708 );
3709 assert_eq!(
3710 redacted.matches("Bearer ***").count() + redacted.matches("bearer ***").count(),
3711 2,
3712 "redacted: {redacted}"
3713 );
3714 assert!(
3715 !redacted.contains("abc.def.ghi") && !redacted.contains("second-token"),
3716 "leaked: {redacted}"
3717 );
3718 }
3719
3720 #[test]
3721 fn redact_proxy_userinfo_strips_password() {
3722 // Corporate-style proxy URL with embedded creds — the
3723 // password must never reach the on-disk log file. URL strings
3724 // are assembled from placeholder constants via `format!` so the
3725 // literal source never contains a scheme-prefixed username +
3726 // password pair (colon-separated, `@`-terminated) that
3727 // GitGuardian's "Basic Auth String" detector would flag as a
3728 // committed credential.
3729 let (placeholder_user, placeholder_pass) = ("PLACEHOLDER_USER", "PLACEHOLDER_PASS");
3730 let with_creds = format!("http://{placeholder_user}:{placeholder_pass}@proxy.example/");
3731 let redacted = redact_proxy_userinfo(&with_creds);
3732 assert_eq!(redacted, "http://***@proxy.example/");
3733 assert!(!redacted.contains(placeholder_pass));
3734 assert!(!redacted.contains(placeholder_user));
3735
3736 // User only (no password) — still redacted.
3737 let with_user_only = format!("https://{placeholder_user}@proxy.example:8080");
3738 let redacted = redact_proxy_userinfo(&with_user_only);
3739 assert_eq!(redacted, "https://***@proxy.example:8080");
3740
3741 // No userinfo segment — pass through.
3742 let redacted = redact_proxy_userinfo("http://proxy.example:3128/");
3743 assert_eq!(redacted, "http://proxy.example:3128/");
3744
3745 // `@` appears only in the path, not as userinfo separator —
3746 // must not be mistaken for credentials.
3747 let redacted = redact_proxy_userinfo("http://proxy.example/path@thing");
3748 assert_eq!(redacted, "http://proxy.example/path@thing");
3749
3750 // Garbage input (no `://`) returned unchanged — the
3751 // surrounding warning log is the only caller and is already
3752 // handling the malformed-URL case.
3753 assert_eq!(redact_proxy_userinfo("not-a-url"), "not-a-url");
3754 }
3755
3756 #[test]
3757 fn redact_body_preview_masks_api_key_param() {
3758 let redacted = redact_body_preview("error api_key=sk-12345&other=val then TOKEN=second-secret");
3759 assert!(redacted.contains("api_key=***"), "redacted: {redacted}");
3760 assert!(redacted.contains("TOKEN=***"), "redacted: {redacted}");
3761 assert!(
3762 !redacted.contains("sk-12345") && !redacted.contains("second-secret"),
3763 "leaked: {redacted}"
3764 );
3765 assert!(
3766 redacted.contains("other=val"),
3767 "non-secret preserved: {redacted}"
3768 );
3769 }
3770
3771 #[test]
3772 fn reviewed_plugin_server_errors_suppress_arbitrary_details() {
3773 let auth = McpHttpAuth {
3774 suppress_server_error_details: true,
3775 ..Default::default()
3776 };
3777 assert_eq!(
3778 auth.server_error_preview("arbitrary credential value"),
3779 "<server details suppressed for reviewed plugin>"
3780 );
3781
3782 let response = serde_json::json!({
3783 "error": { "message": "arbitrary credential value" }
3784 });
3785 let error = response_result(&response, "tools/call", true)
3786 .expect_err("reviewed plugin JSON-RPC error must be generic")
3787 .to_string();
3788 assert!(!error.contains("arbitrary credential value"));
3789 assert!(error.contains("details suppressed"));
3790 }
3791
3792 #[test]
3793 fn invalid_json_preview_collapses_lines_and_redacts_secrets() {
3794 let preview = invalid_json_preview(
3795 b"Authorization: Bearer PLACEHOLDER_TOKEN\nAllow connection? api_key=PLACEHOLDER_KEY",
3796 );
3797
3798 assert!(
3799 preview.contains("Authorization: Bearer *** Allow connection? api_key=***"),
3800 "preview: {preview}"
3801 );
3802 assert!(
3803 !preview.contains('\n'),
3804 "preview should be single-line: {preview}"
3805 );
3806 assert!(
3807 !preview.contains("PLACEHOLDER_TOKEN") && !preview.contains("PLACEHOLDER_KEY"),
3808 "secret leaked: {preview}"
3809 );
3810 }
3811
3812 /// #420: `StdioTransport::shutdown` reaps the child process by sending
3813 /// SIGTERM and giving it a brief grace period before drop fires SIGKILL.
3814 /// The test spawns `cat` (which exits immediately on stdin EOF / SIGTERM)
3815 /// and verifies the transport tears down cleanly. Unix-only because
3816 /// SIGTERM doesn't exist on Windows; on Windows the test would just
3817 /// duplicate the kill_on_drop path.
3818 #[cfg(unix)]
3819 #[tokio::test]
3820 async fn stdio_transport_shutdown_terminates_child() {
3821 use tokio::process::Command as TokioCommand;
3822 let mut cmd = TokioCommand::new("cat");
3823 cmd.stdin(std::process::Stdio::piped())
3824 .stdout(std::process::Stdio::piped())
3825 .stderr(std::process::Stdio::null())
3826 .kill_on_drop(true);
3827 let mut child = cmd.spawn().expect("spawn cat");
3828 let pid = child.id().expect("child pid");
3829 let stdin = child.stdin.take().expect("child stdin");
3830 let stdout = child.stdout.take().expect("child stdout");
3831 let mut transport = StdioTransport {
3832 child: Arc::new(tokio::sync::Mutex::new(child)),
3833 stdin,
3834 reader: tokio::io::BufReader::new(stdout),
3835 stderr_tail: StderrTail::new(),
3836 authority_cancel_watch: None,
3837 _reviewed_launch: None,
3838 };
3839
3840 // shutdown() should send SIGTERM and complete within the grace window.
3841 let start = std::time::Instant::now();
3842 transport.shutdown().await;
3843 let elapsed = start.elapsed();
3844 assert!(
3845 elapsed < STDIO_SHUTDOWN_GRACE + Duration::from_millis(500),
3846 "shutdown blocked beyond grace window: {elapsed:?}"
3847 );
3848
3849 // The child should be reaped — kill(pid, 0) returning ESRCH means
3850 // the pid is gone. If it's still alive, kill(0) returns 0, which
3851 // means our shutdown didn't terminate it.
3852 // SAFETY: pid was just collected from a tokio Child we spawned.
3853 // libc::kill with signal 0 only checks pid existence and is
3854 // async-signal-safe.
3855 let still_alive = unsafe { libc::kill(pid as i32, 0) } == 0;
3856 assert!(
3857 !still_alive,
3858 "child {pid} survived StdioTransport::shutdown — SIGTERM not delivered"
3859 );
3860 }
3861
3862 /// Mid-run MCP server crash: the v0.8.x spawn path used `Stdio::null` for
3863 /// stderr, so a server that died with a useful stderr message left the
3864 /// caller with only "Stdio transport closed". Now stderr is piped into a
3865 /// bounded ring buffer and surfaced when the read side fails.
3866 #[cfg(unix)]
3867 #[tokio::test]
3868 async fn stdio_transport_recv_error_includes_stderr_tail() {
3869 use tokio::process::Command as TokioCommand;
3870
3871 let mut cmd = TokioCommand::new("sh");
3872 cmd.arg("-c")
3873 .arg("echo 'mcp-server: failed to load plugin' 1>&2; exit 1")
3874 .stdin(std::process::Stdio::piped())
3875 .stdout(std::process::Stdio::piped())
3876 .stderr(std::process::Stdio::piped())
3877 .kill_on_drop(true);
3878
3879 let mut child = cmd.spawn().expect("spawn sh");
3880 let stdin = child.stdin.take().expect("stdin");
3881 let stdout = child.stdout.take().expect("stdout");
3882 let stderr = child.stderr.take().expect("stderr");
3883
3884 let stderr_tail = StderrTail::new();
3885 {
3886 let tail = Arc::clone(&stderr_tail);
3887 tokio::spawn(async move {
3888 let mut lines = tokio::io::BufReader::new(stderr).lines();
3889 while let Ok(Some(line)) = lines.next_line().await {
3890 tail.push(line).await;
3891 }
3892 });
3893 }
3894
3895 let mut transport = StdioTransport {
3896 child: Arc::new(tokio::sync::Mutex::new(child)),
3897 stdin,
3898 reader: tokio::io::BufReader::new(stdout),
3899 stderr_tail,
3900 authority_cancel_watch: None,
3901 _reviewed_launch: None,
3902 };
3903
3904 // Give the subprocess time to write its stderr line and exit.
3905 tokio::time::sleep(Duration::from_millis(300)).await;
3906
3907 let err = transport
3908 .recv()
3909 .await
3910 .expect_err("expected transport closed error");
3911 let err_str = format!("{err}");
3912 assert!(
3913 err_str.contains("Stdio transport closed"),
3914 "missing closed marker in: {err_str}"
3915 );
3916 assert!(
3917 err_str.contains("mcp-server: failed to load plugin"),
3918 "stderr context missing from error: {err_str}"
3919 );
3920 }
3921
3922 #[tokio::test]
3923 async fn sse_connect_waits_for_endpoint_before_first_send() {
3924 use std::sync::{
3925 Arc,
3926 atomic::{AtomicBool, Ordering as AtomicOrdering},
3927 };
3928 use tokio::io::{AsyncReadExt, AsyncWriteExt};
3929 use tokio::net::TcpListener;
3930
3931 let _lock = lock_mcp_loopback_tests().await;
3932 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3933 let addr = listener.local_addr().unwrap();
3934 let post_seen = Arc::new(AtomicBool::new(false));
3935 let server_post_seen = Arc::clone(&post_seen);
3936 let cancel_token = tokio_util::sync::CancellationToken::new();
3937 let server_cancel = cancel_token.clone();
3938
3939 let server = tokio::spawn(async move {
3940 loop {
3941 let Ok((mut socket, _)) = listener.accept().await else {
3942 break;
3943 };
3944 let post_seen = Arc::clone(&server_post_seen);
3945 let server_cancel = server_cancel.clone();
3946 tokio::spawn(async move {
3947 let mut request = Vec::new();
3948 let mut buf = [0; 1024];
3949 loop {
3950 let n = socket.read(&mut buf).await.unwrap();
3951 if n == 0 {
3952 return;
3953 }
3954 request.extend_from_slice(&buf[..n]);
3955 if request.windows(4).any(|window| window == b"\r\n\r\n") {
3956 break;
3957 }
3958 }
3959 let request = String::from_utf8_lossy(&request);
3960 if request.starts_with("GET /sse ") {
3961 socket
3962 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
3963 .await
3964 .unwrap();
3965 tokio::time::sleep(Duration::from_millis(150)).await;
3966 socket
3967 .write_all(b"event: endpoint\ndata: /messages\n\n")
3968 .await
3969 .unwrap();
3970 server_cancel.cancelled().await;
3971 } else if request.starts_with("POST /messages ") {
3972 post_seen.store(true, AtomicOrdering::SeqCst);
3973 socket
3974 .write_all(
3975 b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
3976 )
3977 .await
3978 .unwrap();
3979 }
3980 });
3981 }
3982 });
3983
3984 let client = test_http_client();
3985 let url = format!("http://{addr}/sse");
3986 let mut transport = SseTransport::connect(
3987 client,
3988 url,
3989 McpHttpAuth::default(),
3990 cancel_token.clone(),
3991 Duration::from_secs(2),
3992 )
3993 .await
3994 .unwrap();
3995
3996 transport
3997 .send(json_frame(serde_json::json!({
3998 "jsonrpc": "2.0",
3999 "id": 1,
4000 "method": "initialize"
4001 })))
4002 .await
4003 .unwrap();
4004
4005 assert!(
4006 post_seen.load(AtomicOrdering::SeqCst),
4007 "first SSE send should POST to the discovered endpoint"
4008 );
4009
4010 cancel_token.cancel();
4011 server.abort();
4012 }
4013
4014 #[tokio::test]
4015 async fn sse_connect_accepts_crlf_endpoint_events() {
4016 use std::sync::{
4017 Arc,
4018 atomic::{AtomicBool, Ordering as AtomicOrdering},
4019 };
4020 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4021 use tokio::net::TcpListener;
4022
4023 let _lock = lock_mcp_loopback_tests().await;
4024 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4025 let addr = listener.local_addr().unwrap();
4026 let post_seen = Arc::new(AtomicBool::new(false));
4027 let server_post_seen = Arc::clone(&post_seen);
4028 let cancel_token = tokio_util::sync::CancellationToken::new();
4029 let server_cancel = cancel_token.clone();
4030
4031 let server = tokio::spawn(async move {
4032 loop {
4033 let Ok((mut socket, _)) = listener.accept().await else {
4034 break;
4035 };
4036 let post_seen = Arc::clone(&server_post_seen);
4037 let server_cancel = server_cancel.clone();
4038 tokio::spawn(async move {
4039 let mut request = Vec::new();
4040 let mut buf = [0; 1024];
4041 loop {
4042 let n = socket.read(&mut buf).await.unwrap();
4043 if n == 0 {
4044 return;
4045 }
4046 request.extend_from_slice(&buf[..n]);
4047 if request.windows(4).any(|window| window == b"\r\n\r\n") {
4048 break;
4049 }
4050 }
4051 let request = String::from_utf8_lossy(&request);
4052 if request.starts_with("GET /sse ") {
4053 socket
4054 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
4055 .await
4056 .unwrap();
4057 socket
4058 .write_all(b"event: endpoint\r\ndata: /messages\r\n\r\n")
4059 .await
4060 .unwrap();
4061 server_cancel.cancelled().await;
4062 } else if request.starts_with("POST /messages ") {
4063 post_seen.store(true, AtomicOrdering::SeqCst);
4064 socket
4065 .write_all(
4066 b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
4067 )
4068 .await
4069 .unwrap();
4070 }
4071 });
4072 }
4073 });
4074
4075 let client = test_http_client();
4076 let url = format!("http://{addr}/sse");
4077 let mut transport = SseTransport::connect(
4078 client,
4079 url,
4080 McpHttpAuth::default(),
4081 cancel_token.clone(),
4082 Duration::from_secs(2),
4083 )
4084 .await
4085 .unwrap();
4086
4087 transport
4088 .send(json_frame(serde_json::json!({
4089 "jsonrpc": "2.0",
4090 "id": 1,
4091 "method": "initialize"
4092 })))
4093 .await
4094 .unwrap();
4095
4096 assert!(
4097 post_seen.load(AtomicOrdering::SeqCst),
4098 "first SSE send should POST to the CRLF-discovered endpoint"
4099 );
4100
4101 cancel_token.cancel();
4102 server.abort();
4103 }
4104
4105 #[tokio::test]
4106 async fn sse_transport_applies_custom_headers_to_get_and_post() {
4107 use std::sync::{
4108 Arc,
4109 atomic::{AtomicBool, Ordering as AtomicOrdering},
4110 };
4111 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4112 use tokio::net::TcpListener;
4113
4114 let _lock = lock_mcp_loopback_tests().await;
4115 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4116 let addr = listener.local_addr().unwrap();
4117 let get_header_seen = Arc::new(AtomicBool::new(false));
4118 let post_header_seen = Arc::new(AtomicBool::new(false));
4119 let server_get_header_seen = Arc::clone(&get_header_seen);
4120 let server_post_header_seen = Arc::clone(&post_header_seen);
4121 let cancel_token = tokio_util::sync::CancellationToken::new();
4122 let server_cancel = cancel_token.clone();
4123
4124 let server = tokio::spawn(async move {
4125 loop {
4126 let Ok((mut socket, _)) = listener.accept().await else {
4127 break;
4128 };
4129 let get_header_seen = Arc::clone(&server_get_header_seen);
4130 let post_header_seen = Arc::clone(&server_post_header_seen);
4131 let server_cancel = server_cancel.clone();
4132 tokio::spawn(async move {
4133 let mut request = Vec::new();
4134 let mut buf = [0; 1024];
4135 loop {
4136 let n = socket.read(&mut buf).await.unwrap();
4137 if n == 0 {
4138 return;
4139 }
4140 request.extend_from_slice(&buf[..n]);
4141 if request.windows(4).any(|window| window == b"\r\n\r\n") {
4142 break;
4143 }
4144 }
4145 let request = String::from_utf8_lossy(&request);
4146 let request_lower = request.to_lowercase();
4147 if request.starts_with("GET /sse ") {
4148 if request_lower.contains("x-custom-auth: my-test-token") {
4149 get_header_seen.store(true, AtomicOrdering::SeqCst);
4150 }
4151 socket
4152 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
4153 .await
4154 .unwrap();
4155 socket
4156 .write_all(b"event: endpoint\ndata: /messages\n\n")
4157 .await
4158 .unwrap();
4159 server_cancel.cancelled().await;
4160 } else if request.starts_with("POST /messages ") {
4161 if request_lower.contains("x-custom-auth: my-test-token") {
4162 post_header_seen.store(true, AtomicOrdering::SeqCst);
4163 }
4164 socket
4165 .write_all(
4166 b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
4167 )
4168 .await
4169 .unwrap();
4170 }
4171 });
4172 }
4173 });
4174
4175 let client = test_http_client();
4176 let url = format!("http://{addr}/sse");
4177 let mut headers = HashMap::new();
4178 headers.insert("X-Custom-Auth".to_string(), "my-test-token".to_string());
4179 let mut transport = SseTransport::connect(
4180 client,
4181 url,
4182 McpHttpAuth {
4183 headers,
4184 ..Default::default()
4185 },
4186 cancel_token.clone(),
4187 Duration::from_secs(2),
4188 )
4189 .await
4190 .unwrap();
4191
4192 transport
4193 .send(json_frame(serde_json::json!({
4194 "jsonrpc": "2.0",
4195 "id": 1,
4196 "method": "initialize"
4197 })))
4198 .await
4199 .unwrap();
4200
4201 assert!(
4202 get_header_seen.load(AtomicOrdering::SeqCst),
4203 "legacy SSE GET must include user-configured custom headers"
4204 );
4205 assert!(
4206 post_header_seen.load(AtomicOrdering::SeqCst),
4207 "legacy SSE POST must include user-configured custom headers"
4208 );
4209
4210 cancel_token.cancel();
4211 server.abort();
4212 }
4213
4214 #[tokio::test]
4215 async fn sse_post_error_includes_response_body_excerpt() {
4216 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4217 use tokio::net::TcpListener;
4218
4219 let _lock = lock_mcp_loopback_tests().await;
4220 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4221 let addr = listener.local_addr().unwrap();
4222 let cancel_token = tokio_util::sync::CancellationToken::new();
4223 let server_cancel = cancel_token.clone();
4224
4225 let server = tokio::spawn(async move {
4226 loop {
4227 let Ok((mut socket, _)) = listener.accept().await else {
4228 break;
4229 };
4230 let server_cancel = server_cancel.clone();
4231 tokio::spawn(async move {
4232 let mut request = Vec::new();
4233 let mut buf = [0; 1024];
4234 loop {
4235 let n = socket.read(&mut buf).await.unwrap();
4236 if n == 0 {
4237 return;
4238 }
4239 request.extend_from_slice(&buf[..n]);
4240 if request.windows(4).any(|window| window == b"\r\n\r\n") {
4241 break;
4242 }
4243 }
4244 let request = String::from_utf8_lossy(&request);
4245 if request.starts_with("GET /sse ") {
4246 socket
4247 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
4248 .await
4249 .unwrap();
4250 socket
4251 .write_all(b"event: endpoint\ndata: /messages\n\n")
4252 .await
4253 .unwrap();
4254 server_cancel.cancelled().await;
4255 } else if request.starts_with("POST /messages ") {
4256 socket
4257 .write_all(
4258 b"HTTP/1.1 400 Bad Request\r\nConnection: close\r\nContent-Type: application/json\r\nContent-Length: 25\r\n\r\n{\"error\":\"missing query\"}",
4259 )
4260 .await
4261 .unwrap();
4262 }
4263 });
4264 }
4265 });
4266
4267 let client = test_http_client();
4268 let url = format!("http://{addr}/sse");
4269 let mut transport = SseTransport::connect(
4270 client,
4271 url,
4272 McpHttpAuth::default(),
4273 cancel_token.clone(),
4274 Duration::from_secs(2),
4275 )
4276 .await
4277 .unwrap();
4278
4279 let err = transport
4280 .send(json_frame(serde_json::json!({
4281 "jsonrpc": "2.0",
4282 "id": 1,
4283 "method": "initialize"
4284 })))
4285 .await
4286 .expect_err("POST rejection should be returned");
4287 let err = format!("{err:#}");
4288 assert!(
4289 err.contains("400 Bad Request") && err.contains("missing query"),
4290 "SSE POST error should include status and body, got: {err}"
4291 );
4292
4293 cancel_token.cancel();
4294 server.abort();
4295 }
4296
4297 #[tokio::test]
4298 async fn streamable_http_caps_chunked_bodies_without_content_length() {
4299 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4300 use tokio::net::TcpListener;
4301
4302 let _lock = lock_mcp_loopback_tests().await;
4303 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4304 let addr = listener.local_addr().unwrap();
4305
4306 // Serve chunked responses (no Content-Length) of the requested size:
4307 // GET /over streams past the cap, GET /under stays below it.
4308 let server = tokio::spawn(async move {
4309 loop {
4310 let Ok((mut socket, _)) = listener.accept().await else {
4311 break;
4312 };
4313 tokio::spawn(async move {
4314 let mut request = Vec::new();
4315 let mut buf = [0; 1024];
4316 loop {
4317 let n = socket.read(&mut buf).await.unwrap();
4318 if n == 0 {
4319 return;
4320 }
4321 request.extend_from_slice(&buf[..n]);
4322 if request.windows(4).any(|window| window == b"\r\n\r\n") {
4323 break;
4324 }
4325 }
4326 let request = String::from_utf8_lossy(&request);
4327 let total: usize = if request.starts_with("GET /over ") {
4328 256
4329 } else {
4330 16
4331 };
4332 socket
4333 .write_all(
4334 b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nTransfer-Encoding: chunked\r\n\r\n",
4335 )
4336 .await
4337 .unwrap();
4338 let chunk = [b'x'; 32];
4339 let mut sent = 0;
4340 while sent < total {
4341 let n = chunk.len().min(total - sent);
4342 let frame = format!("{n:x}\r\n");
4343 socket.write_all(frame.as_bytes()).await.unwrap();
4344 socket.write_all(&chunk[..n]).await.unwrap();
4345 socket.write_all(b"\r\n").await.unwrap();
4346 sent += n;
4347 }
4348 socket.write_all(b"0\r\n\r\n").await.unwrap();
4349 socket.flush().await.unwrap();
4350 });
4351 }
4352 });
4353
4354 let client = test_http_client();
4355 let cap = 64;
4356
4357 let over = client
4358 .get(format!("http://{addr}/over"))
4359 .send()
4360 .await
4361 .unwrap();
4362 assert_eq!(
4363 over.content_length(),
4364 None,
4365 "chunked response must not declare a length for this test to be meaningful"
4366 );
4367 let err = streamable_http::read_body_capped(over, cap)
4368 .await
4369 .expect_err("a chunked body past the cap must fail, not OOM");
4370 assert!(
4371 err.to_string().contains("exceeds"),
4372 "unexpected error: {err}"
4373 );
4374
4375 let under = client
4376 .get(format!("http://{addr}/under"))
4377 .send()
4378 .await
4379 .unwrap();
4380 let body = streamable_http::read_body_capped(under, cap)
4381 .await
4382 .expect("a chunked body under the cap reads fine");
4383 assert_eq!(body, "x".repeat(16));
4384
4385 server.abort();
4386 }
4387
4388 #[tokio::test]
4389 async fn error_body_excerpt_stops_at_cap_without_waiting_for_eof() {
4390 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4391 use tokio::net::TcpListener;
4392
4393 let _lock = lock_mcp_loopback_tests().await;
4394 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4395 let addr = listener.local_addr().unwrap();
4396 let server_cancel = tokio_util::sync::CancellationToken::new();
4397 let task_cancel = server_cancel.clone();
4398
4399 // Deliberately omit the terminating zero-sized chunk and keep the socket
4400 // open. A `.text()`-based diagnostic would wait for EOF; the bounded
4401 // reader must return as soon as the first chunk reaches the cap.
4402 let server = tokio::spawn(async move {
4403 let (mut socket, _) = listener.accept().await.unwrap();
4404 let mut request = Vec::new();
4405 let mut buf = [0; 1024];
4406 loop {
4407 let n = socket.read(&mut buf).await.unwrap();
4408 if n == 0 {
4409 return;
4410 }
4411 request.extend_from_slice(&buf[..n]);
4412 if request.windows(4).any(|window| window == b"\r\n\r\n") {
4413 break;
4414 }
4415 }
4416 socket
4417 .write_all(
4418 b"HTTP/1.1 500 Internal Server Error\r\nContent-Type: text/plain\r\nTransfer-Encoding: chunked\r\n\r\n100\r\n",
4419 )
4420 .await
4421 .unwrap();
4422 socket.write_all(&[b'x'; 256]).await.unwrap();
4423 socket.write_all(b"\r\n").await.unwrap();
4424 socket.flush().await.unwrap();
4425 task_cancel.cancelled().await;
4426 });
4427
4428 let response = test_http_client()
4429 .get(format!("http://{addr}/preview"))
4430 .send()
4431 .await
4432 .unwrap();
4433 let preview = tokio::time::timeout(Duration::from_secs(1), bounded_body_excerpt(response, 64))
4434 .await
4435 .expect("bounded excerpt must not wait for an attacker-controlled EOF");
4436 assert_eq!(preview, format!("{}…", "x".repeat(64)));
4437
4438 server_cancel.cancel();
4439 server.abort();
4440 }
4441
4442 #[tokio::test]
4443 async fn streamable_http_stale_session_reconnects_and_retries_tool_call() {
4444 use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
4445 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4446 use tokio::net::TcpListener;
4447
4448 async fn write_response(socket: &mut tokio::net::TcpStream, response: &[u8]) {
4449 socket.write_all(response).await.unwrap();
4450 socket.flush().await.unwrap();
4451 socket.shutdown().await.unwrap();
4452 }
4453
4454 let _lock = lock_mcp_loopback_tests().await;
4455 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4456 let addr = listener.local_addr().unwrap();
4457 let get_count = Arc::new(AtomicUsize::new(0));
4458 let stale_seen = Arc::new(AtomicBool::new(false));
4459 let success_seen = Arc::new(AtomicBool::new(false));
4460 let server_get_count = Arc::clone(&get_count);
4461 let server_stale_seen = Arc::clone(&stale_seen);
4462 let server_success_seen = Arc::clone(&success_seen);
4463
4464 let server = tokio::spawn(async move {
4465 loop {
4466 let Ok((mut socket, _)) = listener.accept().await else {
4467 break;
4468 };
4469 let get_count = Arc::clone(&server_get_count);
4470 let stale_seen = Arc::clone(&server_stale_seen);
4471 let success_seen = Arc::clone(&server_success_seen);
4472 tokio::spawn(async move {
4473 let mut request = Vec::new();
4474 let mut buf = [0; 4096];
4475 let header_end = loop {
4476 let n = socket.read(&mut buf).await.unwrap();
4477 if n == 0 {
4478 return;
4479 }
4480 request.extend_from_slice(&buf[..n]);
4481 if let Some(pos) = request.windows(4).position(|w| w == b"\r\n\r\n") {
4482 break pos + 4;
4483 }
4484 };
4485 let headers = String::from_utf8_lossy(&request[..header_end]).to_string();
4486 let content_length = headers
4487 .lines()
4488 .find_map(|line| {
4489 let (name, value) = line.split_once(':')?;
4490 name.eq_ignore_ascii_case("content-length")
4491 .then(|| value.trim().parse::<usize>().ok())
4492 .flatten()
4493 })
4494 .unwrap_or(0);
4495 while request.len() < header_end + content_length {
4496 let n = socket.read(&mut buf).await.unwrap();
4497 if n == 0 {
4498 return;
4499 }
4500 request.extend_from_slice(&buf[..n]);
4501 }
4502 let body = &request[header_end..header_end + content_length];
4503 let session_header = headers.lines().find_map(|line| {
4504 let (name, value) = line.split_once(':')?;
4505 name.eq_ignore_ascii_case("mcp-session-id")
4506 .then(|| value.trim().to_string())
4507 });
4508
4509 if headers.starts_with("GET /mcp ") {
4510 let count = get_count.fetch_add(1, AtomicOrdering::SeqCst);
4511 let session = if count == 0 { "sess-old" } else { "sess-new" };
4512 let response = format!(
4513 "HTTP/1.1 200 OK\r\nConnection: close\r\nMcp-Session-Id: {session}\r\nContent-Length: 0\r\n\r\n"
4514 );
4515 write_response(&mut socket, response.as_bytes()).await;
4516 return;
4517 }
4518
4519 let request_json: serde_json::Value = serde_json::from_slice(body).unwrap();
4520 let method = request_json
4521 .get("method")
4522 .and_then(serde_json::Value::as_str)
4523 .unwrap_or("");
4524 let id = request_json
4525 .get("id")
4526 .cloned()
4527 .unwrap_or_else(|| serde_json::json!("0"));
4528
4529 if method == "tools/call" && session_header.as_deref() == Some("sess-old") {
4530 stale_seen.store(true, AtomicOrdering::SeqCst);
4531 write_response(
4532 &mut socket,
4533 b"HTTP/1.1 404 Not Found\r\nConnection: close\r\nContent-Type: application/json\r\nContent-Length: 27\r\n\r\n{\"error\":\"session expired\"}",
4534 )
4535 .await;
4536 return;
4537 }
4538
4539 let result = match method {
4540 "initialize" => serde_json::json!({
4541 "protocolVersion": "2024-11-05",
4542 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
4543 }),
4544 "tools/list" => serde_json::json!({
4545 "tools": [
4546 { "name": "search", "inputSchema": {} }
4547 ]
4548 }),
4549 "resources/list" => serde_json::json!({ "resources": [] }),
4550 "resources/templates/list" => {
4551 serde_json::json!({ "resourceTemplates": [] })
4552 }
4553 "prompts/list" => serde_json::json!({ "prompts": [] }),
4554 "tools/call" => {
4555 assert_eq!(session_header.as_deref(), Some("sess-new"));
4556 success_seen.store(true, AtomicOrdering::SeqCst);
4557 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
4558 }
4559 _ => {
4560 write_response(
4561 &mut socket,
4562 b"HTTP/1.1 202 Accepted\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
4563 )
4564 .await;
4565 return;
4566 }
4567 };
4568 let response_body = serde_json::json!({
4569 "jsonrpc": "2.0",
4570 "id": id,
4571 "result": result
4572 })
4573 .to_string();
4574 let response = format!(
4575 "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}",
4576 response_body.len(),
4577 response_body
4578 );
4579 write_response(&mut socket, response.as_bytes()).await;
4580 });
4581 }
4582 });
4583
4584 let mut cfg = McpConfig::default();
4585 cfg.servers.insert(
4586 "dephy".to_string(),
4587 McpServerConfig {
4588 command: None,
4589 args: Vec::new(),
4590 env: HashMap::new(),
4591 cwd: None,
4592 url: Some(format!("http://{addr}/mcp")),
4593 transport: None,
4594 connect_timeout: Some(10),
4595 execute_timeout: Some(10),
4596 read_timeout: None,
4597 disabled: false,
4598 enabled: true,
4599 required: false,
4600 enabled_tools: Vec::new(),
4601 disabled_tools: Vec::new(),
4602 headers: HashMap::new(),
4603 env_headers: HashMap::new(),
4604 bearer_token_env_var: None,
4605 scopes: Vec::new(),
4606 oauth: None,
4607 oauth_resource: None,
4608 reviewed_plugin: None,
4609 },
4610 );
4611 let mut pool = McpPool::new(cfg);
4612
4613 let result = pool
4614 .call_tool("mcp_dephy_search", serde_json::json!({ "query": "dephy" }))
4615 .await
4616 .unwrap();
4617
4618 assert_eq!(
4619 result,
4620 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
4621 );
4622 assert!(stale_seen.load(AtomicOrdering::SeqCst));
4623 assert!(success_seen.load(AtomicOrdering::SeqCst));
4624 assert_eq!(get_count.load(AtomicOrdering::SeqCst), 2);
4625
4626 server.abort();
4627 }
4628
4629 #[tokio::test]
4630 async fn legacy_sse_session_expiry_is_marked_stale() {
4631 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4632 use tokio::net::TcpListener;
4633 use tokio::sync::mpsc;
4634
4635 let _lock = lock_mcp_loopback_tests().await;
4636 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4637 let addr = listener.local_addr().unwrap();
4638
4639 let server = tokio::spawn(async move {
4640 let (mut socket, _) = listener.accept().await.unwrap();
4641 let mut request = Vec::new();
4642 let mut buf = [0; 4096];
4643 let header_end = loop {
4644 let n = socket.read(&mut buf).await.unwrap();
4645 if n == 0 {
4646 return;
4647 }
4648 request.extend_from_slice(&buf[..n]);
4649 if let Some(pos) = request.windows(4).position(|w| w == b"\r\n\r\n") {
4650 break pos + 4;
4651 }
4652 };
4653 let headers = String::from_utf8_lossy(&request[..header_end]);
4654 assert!(headers.starts_with("POST /messages "));
4655 socket
4656 .write_all(
4657 b"HTTP/1.1 400 Bad Request\r\nConnection: close\r\nContent-Type: application/json\r\nContent-Length: 27\r\n\r\n{\"error\":\"session expired\"}",
4658 )
4659 .await
4660 .unwrap();
4661 });
4662
4663 let (_sender, receiver) = mpsc::channel(1);
4664 let sse_task = tokio::spawn(async {});
4665 let mut transport = SseTransport {
4666 client: test_http_client(),
4667 base_url: format!("http://{addr}/sse"),
4668 auth: McpHttpAuth::default(),
4669 endpoint_url: Some(format!("http://{addr}/messages")),
4670 receiver,
4671 sse_task,
4672 };
4673
4674 let err = transport
4675 .send(br#"{"jsonrpc":"2.0","id":1,"method":"tools/call"}"#.to_vec())
4676 .await
4677 .expect_err("expired SSE session should fail");
4678
4679 assert!(
4680 is_mcp_stale_session_error(&err),
4681 "SSE session expiry should be retryable, got: {err:#}"
4682 );
4683
4684 server.abort();
4685 }
4686
4687 #[tokio::test]
4688 async fn legacy_sse_closed_stream_reconnects_and_retries_tool_call() {
4689 use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
4690 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4691 use tokio::net::{TcpListener, TcpStream};
4692 use tokio::sync::mpsc;
4693
4694 async fn read_http_request(socket: &mut TcpStream) -> (String, serde_json::Value) {
4695 let mut request = Vec::new();
4696 let mut buf = [0; 4096];
4697 let header_end = loop {
4698 let n = socket.read(&mut buf).await.unwrap();
4699 if n == 0 {
4700 return (String::new(), serde_json::Value::Null);
4701 }
4702 request.extend_from_slice(&buf[..n]);
4703 if let Some(pos) = request.windows(4).position(|w| w == b"\r\n\r\n") {
4704 break pos + 4;
4705 }
4706 };
4707 let headers = String::from_utf8_lossy(&request[..header_end]).to_string();
4708 let content_length = headers
4709 .lines()
4710 .find_map(|line| {
4711 let (name, value) = line.split_once(':')?;
4712 name.eq_ignore_ascii_case("content-length")
4713 .then(|| value.trim().parse::<usize>().ok())
4714 .flatten()
4715 })
4716 .unwrap_or(0);
4717 while request.len() < header_end + content_length {
4718 let n = socket.read(&mut buf).await.unwrap();
4719 if n == 0 {
4720 return (headers, serde_json::Value::Null);
4721 }
4722 request.extend_from_slice(&buf[..n]);
4723 }
4724 let body = &request[header_end..header_end + content_length];
4725 let json = if body.is_empty() {
4726 serde_json::Value::Null
4727 } else {
4728 serde_json::from_slice(body).unwrap()
4729 };
4730 (headers, json)
4731 }
4732
4733 let _lock = lock_mcp_loopback_tests().await;
4734 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4735 let addr = listener.local_addr().unwrap();
4736 let active_sse = Arc::new(Mutex::new(None::<mpsc::UnboundedSender<Option<String>>>));
4737 let get_count = Arc::new(AtomicUsize::new(0));
4738 let tool_call_count = Arc::new(AtomicUsize::new(0));
4739 let success_seen = Arc::new(AtomicBool::new(false));
4740 let server_active_sse = Arc::clone(&active_sse);
4741 let server_get_count = Arc::clone(&get_count);
4742 let server_tool_call_count = Arc::clone(&tool_call_count);
4743 let server_success_seen = Arc::clone(&success_seen);
4744
4745 let server = tokio::spawn(async move {
4746 loop {
4747 let Ok((mut socket, _)) = listener.accept().await else {
4748 break;
4749 };
4750 let active_sse = Arc::clone(&server_active_sse);
4751 let get_count = Arc::clone(&server_get_count);
4752 let tool_call_count = Arc::clone(&server_tool_call_count);
4753 let success_seen = Arc::clone(&server_success_seen);
4754 tokio::spawn(async move {
4755 let (headers, request_json) = read_http_request(&mut socket).await;
4756 if headers.starts_with("GET /sse ") {
4757 get_count.fetch_add(1, AtomicOrdering::SeqCst);
4758 let (tx, mut rx) = mpsc::unbounded_channel::<Option<String>>();
4759 *active_sse.lock().unwrap() = Some(tx);
4760 socket
4761 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
4762 .await
4763 .unwrap();
4764 socket
4765 .write_all(b"event: endpoint\ndata: /messages\n\n")
4766 .await
4767 .unwrap();
4768 while let Some(message) = rx.recv().await {
4769 let Some(message) = message else {
4770 return;
4771 };
4772 let event = format!("event: message\ndata: {message}\n\n");
4773 socket.write_all(event.as_bytes()).await.unwrap();
4774 }
4775 return;
4776 }
4777
4778 if !headers.starts_with("POST /messages ") {
4779 return;
4780 }
4781
4782 socket
4783 .write_all(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
4784 .await
4785 .unwrap();
4786
4787 let method = request_json
4788 .get("method")
4789 .and_then(serde_json::Value::as_str)
4790 .unwrap_or("");
4791 if method == "notifications/initialized" {
4792 return;
4793 }
4794
4795 let id = request_json
4796 .get("id")
4797 .cloned()
4798 .unwrap_or_else(|| serde_json::json!("0"));
4799
4800 if method == "tools/call" {
4801 let count = tool_call_count.fetch_add(1, AtomicOrdering::SeqCst);
4802 if count == 0 {
4803 if let Some(tx) = active_sse.lock().unwrap().take() {
4804 let _ = tx.send(None);
4805 }
4806 return;
4807 }
4808 }
4809
4810 let result = match method {
4811 "initialize" => serde_json::json!({
4812 "protocolVersion": "2024-11-05",
4813 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
4814 }),
4815 "tools/list" => serde_json::json!({
4816 "tools": [
4817 { "name": "search", "inputSchema": {} }
4818 ]
4819 }),
4820 "resources/list" => serde_json::json!({ "resources": [] }),
4821 "resources/templates/list" => {
4822 serde_json::json!({ "resourceTemplates": [] })
4823 }
4824 "prompts/list" => serde_json::json!({ "prompts": [] }),
4825 "tools/call" => {
4826 success_seen.store(true, AtomicOrdering::SeqCst);
4827 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
4828 }
4829 other => panic!("unexpected method: {other}"),
4830 };
4831 let response = serde_json::json!({
4832 "jsonrpc": "2.0",
4833 "id": id,
4834 "result": result
4835 })
4836 .to_string();
4837 // Deliver the response over the *current* SSE channel. The
4838 // retry tool call can race ahead of the reconnecting GET
4839 // /sse that re-stores the sender; under parallel load those
4840 // two server tasks are scheduled in either order, so wait
4841 // briefly for the channel instead of dropping the response
4842 // (which left the client hanging until timeout) (#2597).
4843 let send_deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
4844 let tx = loop {
4845 if let Some(tx) = active_sse.lock().unwrap().as_ref().cloned() {
4846 break Some(tx);
4847 }
4848 if std::time::Instant::now() >= send_deadline {
4849 break None;
4850 }
4851 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
4852 };
4853 if let Some(tx) = tx {
4854 let _ = tx.send(Some(response));
4855 }
4856 });
4857 }
4858 });
4859
4860 let mut cfg = McpConfig::default();
4861 cfg.servers.insert(
4862 "dephy".to_string(),
4863 McpServerConfig {
4864 command: None,
4865 args: Vec::new(),
4866 env: HashMap::new(),
4867 cwd: None,
4868 url: Some(format!("http://{addr}/sse")),
4869 transport: Some("sse".to_string()),
4870 connect_timeout: Some(10),
4871 execute_timeout: Some(10),
4872 read_timeout: None,
4873 disabled: false,
4874 enabled: true,
4875 required: false,
4876 enabled_tools: Vec::new(),
4877 disabled_tools: Vec::new(),
4878 headers: HashMap::new(),
4879 env_headers: HashMap::new(),
4880 bearer_token_env_var: None,
4881 scopes: Vec::new(),
4882 oauth: None,
4883 oauth_resource: None,
4884 reviewed_plugin: None,
4885 },
4886 );
4887 let mut pool = McpPool::new(cfg);
4888
4889 let result = pool
4890 .call_tool("mcp_dephy_search", serde_json::json!({ "query": "dephy" }))
4891 .await
4892 .unwrap();
4893
4894 assert_eq!(
4895 result,
4896 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
4897 );
4898 assert_eq!(tool_call_count.load(AtomicOrdering::SeqCst), 2);
4899 assert_eq!(get_count.load(AtomicOrdering::SeqCst), 2);
4900 assert!(success_seen.load(AtomicOrdering::SeqCst));
4901
4902 server.abort();
4903 }
4904
4905 #[test]
4906 fn session_id_starts_none() {
4907 let transport = StreamableHttpTransport::new(
4908 test_http_client(),
4909 "https://example.invalid/mcp".to_string(),
4910 McpHttpAuth::default(),
4911 );
4912 assert!(transport.session_id.is_none());
4913 }
4914
4915 /// Session ID captured from a POST response is replayed on the next POST.
4916 #[tokio::test]
4917 async fn session_id_captured_from_post_response_and_replayed() {
4918 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4919 use tokio::net::TcpListener;
4920
4921 let _lock = lock_mcp_loopback_tests().await;
4922 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4923 let addr = listener.local_addr().unwrap();
4924 let server = tokio::spawn(async move {
4925 let (mut socket, _) = listener.accept().await.unwrap();
4926 let mut buf = [0u8; 4096];
4927 let n = socket.read(&mut buf).await.unwrap();
4928 let req = String::from_utf8_lossy(&buf[..n]);
4929 assert!(req.starts_with("POST "), "expected POST, got: {req}");
4930
4931 // First POST: return a session ID so the transport captures it.
4932 socket
4933 .write_all(
4934 b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nMcp-Session-Id: sess-abc-123\r\nContent-Length: 2\r\n\r\n{}",
4935 )
4936 .await
4937 .unwrap();
4938 socket.flush().await.unwrap();
4939
4940 // Read the second POST — should contain the session ID.
4941 let mut buf2 = [0u8; 4096];
4942 let n2 = socket.read(&mut buf2).await.unwrap();
4943 let req2 = String::from_utf8_lossy(&buf2[..n2]);
4944 // reqwest lower-cases header names.
4945 let req2_lower = req2.to_lowercase();
4946 assert!(
4947 req2_lower.contains("mcp-session-id: sess-abc-123"),
4948 "second POST must replay captured session ID, got:\n{req2}"
4949 );
4950
4951 socket
4952 .write_all(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
4953 .await
4954 .unwrap();
4955 });
4956
4957 let client = test_http_client();
4958 let url = format!("http://{addr}/mcp");
4959 let mut transport = StreamableHttpTransport::new(client, url, McpHttpAuth::default());
4960
4961 // First send: server returns Mcp-Session-Id.
4962 transport
4963 .send(json_frame(serde_json::json!({
4964 "jsonrpc": "2.0", "id": 1,
4965 "method": "initialize",
4966 "params": {}
4967 })))
4968 .await
4969 .unwrap();
4970 assert_eq!(
4971 transport.session_id.as_deref(),
4972 Some("sess-abc-123"),
4973 "session ID should be captured from response"
4974 );
4975
4976 // Second send: should replay the session ID.
4977 transport
4978 .send(json_frame(serde_json::json!({
4979 "jsonrpc": "2.0", "id": 2,
4980 "method": "tools/list",
4981 "params": {}
4982 })))
4983 .await
4984 .unwrap();
4985
4986 server.abort();
4987 }
4988
4989 /// Custom headers configured in McpServerConfig are applied to the GET
4990 /// preflight so servers that require auth on session-establishment GET
4991 /// (e.g. Hindsight, #1629) can authenticate it.
4992 #[tokio::test]
4993 async fn custom_headers_applied_to_get_preflight() {
4994 use tokio::io::{AsyncReadExt, AsyncWriteExt};
4995 use tokio::net::TcpListener;
4996
4997 let _lock = lock_mcp_loopback_tests().await;
4998 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4999 let addr = listener.local_addr().unwrap();
5000 // The test signals success by writing to this flag — the GET handler
5001 // sets it when it sees the expected header.
5002 let header_seen = Arc::new(AtomicBool::new(false));
5003 let header_seen_srv = Arc::clone(&header_seen);
5004
5005 let server = tokio::spawn(async move {
5006 let (mut socket, _) = listener.accept().await.unwrap();
5007 let mut buf = [0u8; 4096];
5008 let n = socket.read(&mut buf).await.unwrap();
5009 let req = String::from_utf8_lossy(&buf[..n]);
5010
5011 // reqwest lower-cases header names.
5012 if req.starts_with("GET ") && req.to_lowercase().contains("x-custom-auth: my-test-token") {
5013 header_seen_srv.store(true, AtomicOrdering::SeqCst);
5014 }
5015
5016 socket
5017 .write_all(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
5018 .await
5019 .unwrap();
5020 });
5021
5022 let client = test_http_client();
5023 let url = format!("http://{addr}/mcp");
5024 let mut headers = HashMap::new();
5025 headers.insert("X-Custom-Auth".to_string(), "my-test-token".to_string());
5026
5027 let mut transport = HttpTransport::new(
5028 client,
5029 url,
5030 McpHttpAuth {
5031 headers,
5032 ..Default::default()
5033 },
5034 tokio_util::sync::CancellationToken::new(),
5035 Duration::from_secs(10),
5036 );
5037
5038 transport.try_establish_session().await.unwrap();
5039
5040 server.abort();
5041
5042 assert!(
5043 header_seen.load(AtomicOrdering::SeqCst),
5044 "GET preflight must include user-configured custom headers"
5045 );
5046 }
5047
5048 // === add_runtime_server_config conflict tests ===
5049
5050 #[test]
5051 fn add_runtime_server_config_rejects_static_conflict() {
5052 let config: McpConfig = serde_json::from_str(
5053 r#"{
5054 "servers": {
5055 "existing": {"command": "node server.js"}
5056 }
5057 }"#,
5058 )
5059 .unwrap();
5060 let pool = McpPool::new(config);
5061
5062 let err = pool
5063 .add_runtime_server_config(
5064 "existing".to_string(),
5065 serde_json::from_str(r#"{"command": "npx other"}"#).unwrap(),
5066 )
5067 .unwrap_err();
5068 assert!(err.contains("already exists in the config file"));
5069 }
5070
5071 #[test]
5072 fn add_runtime_server_config_rejects_dynamic_duplicate() {
5073 let pool = McpPool::new(McpConfig::default());
5074
5075 pool.add_runtime_server_config(
5076 "my_server".to_string(),
5077 serde_json::from_str(r#"{"command": "node a.js"}"#).unwrap(),
5078 )
5079 .unwrap();
5080
5081 let err = pool
5082 .add_runtime_server_config(
5083 "my_server".to_string(),
5084 serde_json::from_str(r#"{"command": "node b.js"}"#).unwrap(),
5085 )
5086 .unwrap_err();
5087 assert!(err.contains("already started earlier"));
5088 }
5089
5090 #[test]
5091 fn add_runtime_server_config_accepts_new_name() {
5092 let pool = McpPool::new(McpConfig::default());
5093
5094 pool.add_runtime_server_config(
5095 "brand_new".to_string(),
5096 serde_json::from_str(r#"{"command": "node x.js"}"#).unwrap(),
5097 )
5098 .unwrap();
5099 }
5100
5101 /// Server attribution and the model-facing tool name must come from one
5102 /// definition. If they ever drift, a human reading tool provenance would be
5103 /// told which server owns a name the model never saw.
5104 #[test]
5105 fn mcp_model_tool_names_and_server_attribution_share_one_definition() {
5106 assert_eq!(
5107 McpPool::mcp_model_tool_name("files", "read"),
5108 "mcp_files_read"
5109 );
5110 // A server name containing `_` is exactly why the reverse split is a guess.
5111 assert_eq!(
5112 McpPool::mcp_model_tool_name("my_server", "read_file"),
5113 "mcp_my_server_read_file"
5114 );
5115
5116 let resolved =
5117 McpPool::resolve_tool_server_map([("files", "read"), ("git", "status")].into_iter());
5118 assert_eq!(
5119 resolved.get("mcp_files_read").map(String::as_str),
5120 Some("files")
5121 );
5122 assert_eq!(
5123 resolved.get("mcp_git_status").map(String::as_str),
5124 Some("git")
5125 );
5126
5127 // Ambiguity: two servers collapse onto the same model name. Neither wins,
5128 // so the name resolves to no server and callers report it as unknown —
5129 // the same rule `all_tools` applies when it hides the ambiguous tool.
5130 let ambiguous = McpPool::resolve_tool_server_map(
5131 [("a_b", "c"), ("a", "b_c"), ("solo", "tool")].into_iter(),
5132 );
5133 assert!(
5134 !ambiguous.contains_key("mcp_a_b_c"),
5135 "an ambiguous model name must resolve to no server"
5136 );
5137 assert_eq!(
5138 ambiguous.get("mcp_solo_tool").map(String::as_str),
5139 Some("solo")
5140 );
5141 }
5142
5143 #[test]
5144 fn removed_runtime_server_config_can_be_retried_with_same_name() {
5145 let mut pool = McpPool::new(McpConfig::default());
5146 let config: McpServerConfig = serde_json::from_str(r#"{"command": "node a.js"}"#).unwrap();
5147
5148 pool.add_runtime_server_config("retryable".to_string(), config.clone())
5149 .unwrap();
5150 pool.remove_runtime_server_config("retryable");
5151 pool.add_runtime_server_config("retryable".to_string(), config)
5152 .expect("rollback must release the deterministic runtime name");
5153 }
5154
5154 lines RUST