返回 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::http::HttpTransport;
3 use super::http_client::McpHttpAuth;
4 use super::streamable_http::StreamableHttpTransport;
5 use super::wire::{
6 find_sse_event_separator, find_sse_event_separator_bytes, is_mcp_stale_session_error,
7 parse_sse_message_data,
8 };
9 use super::*;
10 use reqwest::header::{ACCEPT, CONTENT_TYPE};
11 use serde_json::{Value, json};
12 use std::collections::VecDeque;
13 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering as AtomicOrdering};
14 use std::sync::{Arc, Mutex, OnceLock};
15
16 fn test_http_client() -> reqwest::Client {
17 let _ = rustls::crypto::ring::default_provider().install_default();
18 crate::tls::reqwest_client()
19 }
20
21 async fn lock_mcp_loopback_tests() -> tokio::sync::MutexGuard<'static, ()> {
22 static LOCK: OnceLock<tokio::sync::Mutex<()>> = OnceLock::new();
23 LOCK.get_or_init(|| tokio::sync::Mutex::new(()))
24 .lock()
25 .await
26 }
27
28 struct WorkspaceTrustConfigGuard {
29 config_path: PathBuf,
30 _codewhale_config_path: crate::test_support::EnvVarGuard,
31 _deepseek_config_path: crate::test_support::EnvVarGuard,
32 _env_lock: crate::test_support::TestEnvLock,
33 }
34
35 fn workspace_trust_config_guard(workspace: &Path) -> WorkspaceTrustConfigGuard {
36 let env_lock = crate::test_support::lock_test_env();
37 let config_path = workspace
38 .parent()
39 .unwrap_or(workspace)
40 .join("user-config")
41 .join("config.toml");
42 if let Some(parent) = config_path.parent() {
43 fs::create_dir_all(parent).unwrap();
44 }
45 let codewhale_config_path =
46 crate::test_support::EnvVarGuard::set("CODEWHALE_CONFIG_PATH", config_path.as_os_str());
47 let deepseek_config_path = crate::test_support::EnvVarGuard::remove("DEEPSEEK_CONFIG_PATH");
48
49 WorkspaceTrustConfigGuard {
50 config_path,
51 _codewhale_config_path: codewhale_config_path,
52 _deepseek_config_path: deepseek_config_path,
53 _env_lock: env_lock,
54 }
55 }
56
57 fn write_workspace_trust_config(config_path: &Path, workspace: &Path) {
58 let workspace = workspace
59 .canonicalize()
60 .unwrap_or_else(|_| workspace.to_path_buf());
61 let key = workspace
62 .to_string_lossy()
63 .replace('\\', "\\\\")
64 .replace('"', "\\\"");
65 fs::write(
66 config_path,
67 format!("[projects.\"{key}\"]\ntrust_level = \"trusted\"\n"),
68 )
69 .unwrap();
70 }
71
72 fn mark_workspace_trusted(workspace: &Path) -> WorkspaceTrustConfigGuard {
73 let guard = workspace_trust_config_guard(workspace);
74 write_workspace_trust_config(&guard.config_path, workspace);
75 guard
76 }
77
78 #[test]
79 fn test_mcp_config_defaults() {
80 let config = McpConfig::default();
81 assert_eq!(config.timeouts.connect_timeout, 30);
82 assert_eq!(config.timeouts.execute_timeout, 1800);
83 assert_eq!(config.timeouts.read_timeout, 120);
84 assert!(config.servers.is_empty());
85 }
86
87 #[test]
88 fn reviewed_remote_endpoint_identity_normalizes_case_idna_and_default_ports() {
89 let canonical = reviewed_remote_endpoint_identity("https://example.com/mcp").unwrap();
90 assert_eq!(
91 reviewed_remote_endpoint_identity("https://EXAMPLE.COM:443/mcp").unwrap(),
92 canonical
93 );
94 assert_eq!(
95 reviewed_remote_endpoint_identity("https://BÜCHER.example:443/mcp").unwrap(),
96 reviewed_remote_endpoint_identity("https://xn--bcher-kva.example/mcp").unwrap()
97 );
98 assert_ne!(
99 reviewed_remote_endpoint_identity("https://example.com:444/mcp").unwrap(),
100 canonical
101 );
102 assert!(reviewed_remote_endpoint_identity("http://localhost:8080/mcp").is_ok());
103 assert!(reviewed_remote_endpoint_identity("http://127.0.0.1/mcp").is_ok());
104 assert!(reviewed_remote_endpoint_identity("http://[::1]/mcp").is_ok());
105 }
106
107 #[test]
108 fn reviewed_remote_endpoint_identity_rejects_ambiguous_or_secret_bearing_urls() {
109 for endpoint in [
110 "http://example.com/mcp",
111 "ftp://example.com/mcp",
112 "https://user@example.com/mcp",
113 "https://user:secret@example.com/mcp",
114 "https://example.com/mcp?token=secret",
115 "https://example.com/mcp#fragment",
116 ] {
117 let error = reviewed_remote_endpoint_identity(endpoint)
118 .expect_err("unsafe reviewed endpoint must fail closed")
119 .to_string();
120 assert!(
121 !error.contains("secret"),
122 "endpoint error leaked URL material"
123 );
124 }
125 }
126
127 #[test]
128 fn reviewed_plugin_redirects_are_exact_normalized_origin_only() {
129 let approved = reviewed_remote_endpoint_identity("https://BÜCHER.example:443/mcp")
130 .unwrap()
131 .1;
132 let accepted = [
133 "https://xn--bcher-kva.example/next",
134 "https://BÜCHER.example:443/next?cursor=opaque",
135 ];
136 for endpoint in accepted {
137 assert!(reviewed_redirect_matches_origin(
138 &reqwest::Url::parse(endpoint).unwrap(),
139 &approved
140 ));
141 }
142
143 let rejected = [
144 "http://xn--bcher-kva.example/next",
145 "https://user@xn--bcher-kva.example/next",
146 "https://xn--bcher-kva.example:444/next",
147 "https://other.example/next",
148 ];
149 for endpoint in rejected {
150 assert!(!reviewed_redirect_matches_origin(
151 &reqwest::Url::parse(endpoint).unwrap(),
152 &approved
153 ));
154 }
155 }
156
157 #[test]
158 fn reviewed_plugin_remote_proxy_policy_never_reads_ambient_environment() {
159 let reads = std::cell::Cell::new(0_u32);
160 let proxy = configured_mcp_proxy(
161 &reqwest::Url::parse("https://example.com/mcp").unwrap(),
162 true,
163 |_| {
164 reads.set(reads.get() + 1);
165 Ok("http://127.0.0.1:9999".to_string())
166 },
167 );
168
169 assert_eq!(
170 reads.get(),
171 0,
172 "reviewed remotes must not read proxy values"
173 );
174 assert!(proxy.unwrap().is_none());
175 }
176
177 #[test]
178 fn user_authored_mcp_proxy_policy_keeps_environment_support() {
179 let requested = std::cell::RefCell::new(Vec::new());
180 let proxy = configured_mcp_proxy(
181 &reqwest::Url::parse("https://example.com/mcp").unwrap(),
182 false,
183 |name| {
184 requested.borrow_mut().push(name.to_string());
185 match name {
186 "HTTPS_PROXY" => Ok("http://127.0.0.1:8080".to_string()),
187 _ => Err(std::env::VarError::NotPresent),
188 }
189 },
190 );
191
192 assert_eq!(
193 requested.into_inner(),
194 vec![
195 "HTTPS_PROXY".to_string(),
196 "NO_PROXY".to_string(),
197 "no_proxy".to_string(),
198 ]
199 );
200 assert!(proxy.unwrap().is_some());
201 }
202
203 #[test]
204 fn test_mcp_config_parse() {
205 let json = r#"{
206 "timeouts": {
207 "connect_timeout": 15,
208 "execute_timeout": 90
209 },
210 "servers": {
211 "test": {
212 "command": "node",
213 "args": ["server.js"],
214 "env": {"FOO": "bar"}
215 }
216 }
217 }"#;
218
219 let config: McpConfig = serde_json::from_str(json).unwrap();
220 assert_eq!(config.timeouts.connect_timeout, 15);
221 assert_eq!(config.timeouts.execute_timeout, 90);
222 assert_eq!(config.timeouts.read_timeout, 120); // default
223 assert!(config.servers.contains_key("test"));
224
225 let server = config.servers.get("test").unwrap();
226 assert_eq!(server.command, Some("node".to_string()));
227 assert_eq!(server.args, vec!["server.js"]);
228 assert_eq!(server.env.get("FOO"), Some(&"bar".to_string()));
229 }
230
231 #[test]
232 fn mcp_pool_parse_prefixed_name_rejects_ambiguous_configured_server_prefixes() {
233 let config: McpConfig = serde_json::from_str(
234 r#"{
235 "servers": {
236 "my": {"command": "node"},
237 "my_db": {"command": "node"}
238 }
239 }"#,
240 )
241 .unwrap();
242 let pool = McpPool::new(config);
243
244 let error = pool
245 .parse_prefixed_name("mcp_my_db_execute_sql")
246 .expect_err("configured server-prefix collisions must fail closed");
247 assert!(error.to_string().contains("Unknown MCP tool name"));
248 }
249
250 #[test]
251 fn mcp_server_config_parses_custom_headers() {
252 let json = r#"{
253 "servers": {
254 "hf": {
255 "url": "https://example.invalid/mcp",
256 "headers": {
257 "Authorization": "Bearer tok",
258 "X-Org": "anthropic"
259 }
260 }
261 }
262 }"#;
263 let cfg: McpConfig = serde_json::from_str(json).unwrap();
264 let hf = cfg.servers.get("hf").expect("server present");
265 assert_eq!(
266 hf.headers.get("Authorization"),
267 Some(&"Bearer tok".to_string())
268 );
269 assert_eq!(hf.headers.get("X-Org"), Some(&"anthropic".to_string()));
270 }
271
272 #[test]
273 fn mcp_server_config_parses_remote_auth_fields() {
274 let json = r#"{
275 "servers": {
276 "remote": {
277 "url": "https://example.invalid/mcp",
278 "env_http_headers": {
279 "X-Api-Key": "REMOTE_MCP_KEY"
280 },
281 "bearer_token_env_var": "REMOTE_MCP_TOKEN",
282 "scopes": ["tools/read", "tools/write"],
283 "oauth": {
284 "client_id": "client-123"
285 },
286 "oauth_resource": "https://example.invalid"
287 }
288 }
289 }"#;
290 let cfg: McpConfig = serde_json::from_str(json).unwrap();
291 let remote = cfg.servers.get("remote").expect("server present");
292 assert_eq!(
293 remote.env_headers.get("X-Api-Key"),
294 Some(&"REMOTE_MCP_KEY".to_string())
295 );
296 assert_eq!(
297 remote.bearer_token_env_var.as_deref(),
298 Some("REMOTE_MCP_TOKEN")
299 );
300 assert_eq!(remote.scopes, vec!["tools/read", "tools/write"]);
301 assert_eq!(remote.oauth_client_id(), Some("client-123"));
302 assert_eq!(
303 remote.oauth_resource.as_deref(),
304 Some("https://example.invalid")
305 );
306 }
307
308 #[test]
309 fn mcp_server_config_omits_headers_when_empty() {
310 // Empty headers map should not appear in the serialized output —
311 // older mcp.json files written before v0.8.31 must round-trip
312 // unchanged so a `mcp save` from a fresh install doesn't add
313 // dead keys.
314 let cfg = McpServerConfig {
315 command: Some("node".into()),
316 args: vec!["server.js".into()],
317 env: HashMap::new(),
318 cwd: None,
319 url: None,
320 transport: None,
321 connect_timeout: None,
322 execute_timeout: None,
323 read_timeout: None,
324 disabled: false,
325 enabled: true,
326 required: false,
327 enabled_tools: Vec::new(),
328 disabled_tools: Vec::new(),
329 headers: HashMap::new(),
330 env_headers: HashMap::new(),
331 bearer_token_env_var: None,
332 scopes: Vec::new(),
333 oauth: None,
334 oauth_resource: None,
335 reviewed_plugin: None,
336 runtime_added: false,
337 allow_private_network: false,
338 };
339 let serialized = serde_json::to_string(&cfg).unwrap();
340 assert!(
341 !serialized.contains("\"headers\""),
342 "empty headers must be omitted: {serialized}"
343 );
344 assert!(
345 !serialized.contains("\"env_headers\""),
346 "empty env_headers must be omitted: {serialized}"
347 );
348 assert!(
349 !serialized.contains("\"scopes\""),
350 "empty scopes must be omitted: {serialized}"
351 );
352 assert!(
353 !serialized.contains("\"oauth\""),
354 "empty oauth config must be omitted: {serialized}"
355 );
356 }
357
358 #[test]
359 fn expand_env_placeholders_expands_value_from_environment() {
360 let _lock = crate::test_support::lock_test_env();
361 let _secret =
362 crate::test_support::EnvVarGuard::set("MCP_TEST_SECRET_TOKEN", "test-secret-123456");
363 let mut env = HashMap::new();
364 env.insert(
365 "API_TOKEN".to_string(),
366 "${MCP_TEST_SECRET_TOKEN}".to_string(),
367 );
368
369 let expanded = expand_env_placeholders_map(&env, "env").unwrap();
370
371 assert_eq!(
372 expanded.get("API_TOKEN").map(String::as_str),
373 Some("test-secret-123456")
374 );
375 }
376
377 #[test]
378 fn expand_env_placeholders_reports_missing_variable_without_secret_value() {
379 let _lock = crate::test_support::lock_test_env();
380 let _missing = crate::test_support::EnvVarGuard::remove("MCP_TEST_MISSING_SECRET");
381
382 let err = expand_env_placeholders("Bearer ${MCP_TEST_MISSING_SECRET}")
383 .expect_err("missing env should fail")
384 .to_string();
385
386 // The error must name the variable but must not leak the surrounding
387 // value (which in practice carries the secret).
388 assert!(err.contains("MCP_TEST_MISSING_SECRET"));
389 assert!(!err.contains("Bearer "));
390 }
391
392 #[test]
393 fn reviewed_plugin_environment_uses_only_the_pre_dotenv_snapshot() {
394 let _lock = crate::test_support::lock_test_env();
395 let dir = tempfile::tempdir().unwrap();
396 let plugin_base = dir.path().join("plugins/env-snapshot");
397 fs::create_dir_all(&plugin_base).unwrap();
398 fs::write(
399 plugin_base.join("plugin.toml"),
400 "schema_version = 1\n[plugin]\nname = \"env-snapshot\"\nversion = \"1.0.0\"\n",
401 )
402 .unwrap();
403 let (_, authority) = active_plugin_fixture(&plugin_base);
404 let snapshot = crate::plugins::HostEnvironment::from_entries([(
405 OsString::from("PLUGIN_SNAPSHOT_TOKEN"),
406 OsString::from("captured-before-dotenv"),
407 )]);
408 let mut server = test_server_config();
409 server
410 .env
411 .insert("TOKEN".to_string(), "${PLUGIN_SNAPSHOT_TOKEN}".to_string());
412 server.reviewed_plugin =
413 Some(ReviewedPluginMcpSource::from_authority(authority, None, Arc::new(snapshot)).unwrap());
414 let _late_dotenv = crate::test_support::EnvVarGuard::set(
415 "PLUGIN_SNAPSHOT_TOKEN",
416 "workspace-dotenv-must-not-win",
417 );
418
419 let expanded = expanded_mcp_stdio_env(&server).unwrap();
420 assert_eq!(expanded["TOKEN"], "captured-before-dotenv");
421
422 server.reviewed_plugin.as_mut().unwrap().host_environment =
423 Arc::new(crate::plugins::HostEnvironment::from_entries([]));
424 let error = expanded_mcp_stdio_env(&server)
425 .expect_err("a value present only after dotenv must fail closed");
426 assert!(
427 format!("{error:#}").contains("PLUGIN_SNAPSHOT_TOKEN"),
428 "unexpected missing-snapshot error: {error:#}"
429 );
430 assert!(!format!("{error:#}").contains("workspace-dotenv-must-not-win"));
431 }
432
433 fn write_path_only_test_command(dir: &Path) -> String {
434 let command = "codewhale-mcp-path-only-test";
435 #[cfg(windows)]
436 let file_name = format!("{command}.exe");
437 #[cfg(not(windows))]
438 let file_name = command.to_string();
439 let path = dir.join(file_name);
440 fs::write(&path, b"test executable").expect("write path-only test command");
441 #[cfg(unix)]
442 {
443 use std::os::unix::fs::PermissionsExt;
444
445 let mut permissions = fs::metadata(&path)
446 .expect("path-only command metadata")
447 .permissions();
448 permissions.set_mode(0o755);
449 fs::set_permissions(&path, permissions).expect("make path-only test command executable");
450 }
451 command.to_string()
452 }
453
454 #[test]
455 fn static_mcp_command_uses_expanded_sanitized_stdio_path() {
456 let _lock = crate::test_support::lock_test_env();
457 let temp = tempfile::tempdir().expect("tempdir");
458 let command = write_path_only_test_command(temp.path());
459 let _path = crate::test_support::EnvVarGuard::set(
460 "CODEWHALE_MCP_PATH_ONLY_DIR",
461 temp.path().as_os_str(),
462 );
463 let _secret = crate::test_support::EnvVarGuard::set(
464 "CODEWHALE_MCP_STATIC_TEST_SECRET",
465 "must-not-reach-child",
466 );
467 let mut server = test_server_config();
468 server.command = Some(command);
469 server.env.insert(
470 "PATH".to_string(),
471 "${CODEWHALE_MCP_PATH_ONLY_DIR}".to_string(),
472 );
473
474 assert_eq!(
475 static_mcp_command_availability(&server).expect("static command check"),
476 McpCommandAvailability::Available
477 );
478
479 let child_env = mcp_stdio_child_env(&server).expect("stdio child env");
480 assert_eq!(
481 env_value(&child_env, "PATH"),
482 Some(temp.path().as_os_str()),
483 "expanded server PATH must override the inherited PATH"
484 );
485 assert!(
486 child_env
487 .iter()
488 .all(|(key, _)| key != "CODEWHALE_MCP_STATIC_TEST_SECRET"),
489 "static lookup must use the same sanitized parent environment as spawn"
490 );
491
492 let expanded_env = expand_env_placeholders_map(&server.env, "env").expect("expanded env");
493 let mut old_spawn_command = tokio::process::Command::new("unused-test-command");
494 crate::child_env::apply_to_tokio_command_mcp(
495 &mut old_spawn_command,
496 crate::child_env::string_map_env(&expanded_env),
497 );
498 let old_spawn_env = old_spawn_command
499 .as_std()
500 .get_envs()
501 .map(|(key, value)| {
502 (
503 key.to_os_string(),
504 value.expect("spawn env value").to_os_string(),
505 )
506 })
507 .collect::<HashMap<_, _>>();
508 let static_env = child_env.into_iter().collect::<HashMap<_, _>>();
509 assert_eq!(
510 static_env, old_spawn_env,
511 "static lookup and the pre-fix spawn helper must receive identical environments"
512 );
513 }
514
515 #[cfg(not(windows))]
516 #[test]
517 fn static_mcp_command_reports_missing_with_server_path_override() {
518 let temp = tempfile::tempdir().expect("tempdir");
519 let mut server = test_server_config();
520 server.command = Some("codewhale-mcp-command-that-does-not-exist".to_string());
521 server.env.insert(
522 "PATH".to_string(),
523 temp.path().to_string_lossy().into_owned(),
524 );
525
526 assert_eq!(
527 static_mcp_command_availability(&server).expect("static command check"),
528 McpCommandAvailability::Missing
529 );
530 }
531
532 #[test]
533 fn static_mcp_command_reports_invalid_path_expansion() {
534 let _lock = crate::test_support::lock_test_env();
535 let _missing = crate::test_support::EnvVarGuard::remove("CODEWHALE_MCP_MISSING_PATH_DIR");
536 let mut server = test_server_config();
537 server.command = Some("codewhale-mcp-command".to_string());
538 server.env.insert(
539 "PATH".to_string(),
540 "do-not-leak-${CODEWHALE_MCP_MISSING_PATH_DIR}-also-secret".to_string(),
541 );
542
543 let error = static_mcp_command_availability(&server)
544 .expect_err("missing PATH placeholder must fail static validation");
545 let error = format!("{error:#}");
546 assert!(error.contains("CODEWHALE_MCP_MISSING_PATH_DIR"));
547 assert!(!error.contains("codewhale-mcp-command"));
548 assert!(!error.contains("do-not-leak"));
549 assert!(!error.contains("also-secret"));
550 }
551
552 #[cfg(unix)]
553 fn write_unix_test_command(path: &Path, mode: u32) {
554 use std::os::unix::fs::PermissionsExt;
555
556 fs::write(path, b"#!/bin/sh\nexit 0\n").expect("write Unix test command");
557 let mut permissions = fs::metadata(path)
558 .expect("Unix test command metadata")
559 .permissions();
560 permissions.set_mode(mode);
561 fs::set_permissions(path, permissions).expect("set Unix test command mode");
562 }
563
564 #[cfg(unix)]
565 #[test]
566 fn static_mcp_command_anchors_relative_and_empty_path_to_server_cwd() {
567 let temp = tempfile::tempdir().expect("tempdir");
568 let cwd = temp.path().join("server-cwd");
569 let bin = cwd.join("relative-bin");
570 fs::create_dir_all(&bin).expect("relative bin dir");
571 let relative_command = "codewhale-mcp-relative-path-test";
572 write_unix_test_command(&bin.join(relative_command), 0o755);
573
574 let mut server = test_server_config();
575 server.command = Some(relative_command.to_string());
576 server.cwd = Some(cwd.clone());
577 server
578 .env
579 .insert("PATH".to_string(), "relative-bin".to_string());
580 assert_eq!(
581 static_mcp_command_availability(&server).expect("relative PATH check"),
582 McpCommandAvailability::Available
583 );
584
585 let empty_path_command = "codewhale-mcp-empty-path-test";
586 write_unix_test_command(&cwd.join(empty_path_command), 0o755);
587 server.command = Some(empty_path_command.to_string());
588 server.env.insert("PATH".to_string(), String::new());
589 assert_eq!(
590 static_mcp_command_availability(&server).expect("empty PATH check"),
591 McpCommandAvailability::Available,
592 "an empty Unix PATH entry resolves from the child's cwd"
593 );
594 }
595
596 #[cfg(unix)]
597 #[test]
598 fn static_mcp_command_preserves_literal_name_and_requires_execute_bits() {
599 let temp = tempfile::tempdir().expect("tempdir");
600 let literal_command = " codewhale-mcp-literal-command-test ";
601 write_unix_test_command(&temp.path().join(literal_command), 0o755);
602
603 let mut server = test_server_config();
604 server.command = Some(literal_command.to_string());
605 server.env.insert(
606 "PATH".to_string(),
607 temp.path().to_string_lossy().into_owned(),
608 );
609 assert_eq!(
610 static_mcp_command_availability(&server).expect("literal command check"),
611 McpCommandAvailability::Available,
612 "static validation must not trim the command passed to Command::new"
613 );
614
615 let non_executable = temp.path().join("codewhale-mcp-non-executable-test");
616 write_unix_test_command(&non_executable, 0o644);
617 server.command = Some("codewhale-mcp-non-executable-test".to_string());
618 assert_eq!(
619 static_mcp_command_availability(&server).expect("PATH execute-bit check"),
620 McpCommandAvailability::Missing
621 );
622 server.command = Some(non_executable.to_string_lossy().into_owned());
623 assert_eq!(
624 static_mcp_command_availability(&server).expect("absolute execute-bit check"),
625 McpCommandAvailability::Missing
626 );
627 }
628
629 #[cfg(windows)]
630 #[test]
631 fn static_mcp_command_matches_windows_path_and_extension_rules() {
632 let temp = tempfile::tempdir().expect("tempdir");
633 let command = write_path_only_test_command(temp.path());
634 let mut server = test_server_config();
635 server.command = Some(command);
636 server.env.insert(
637 "Path".to_string(),
638 temp.path().to_string_lossy().into_owned(),
639 );
640
641 assert_eq!(
642 static_mcp_command_availability(&server).expect("case-insensitive PATH check"),
643 McpCommandAvailability::Available
644 );
645
646 server.command = Some(
647 temp.path()
648 .join("codewhale-mcp-path-only-test")
649 .to_string_lossy()
650 .into_owned(),
651 );
652 assert_eq!(
653 static_mcp_command_availability(&server).expect("absolute omitted .exe check"),
654 McpCommandAvailability::Available
655 );
656
657 let pathext_command = "codewhale-mcp-pathext-only-test";
658 fs::write(
659 temp.path().join(format!("{pathext_command}.cmd")),
660 b"@exit /b 0\r\n",
661 )
662 .expect("write PATHEXT-only command");
663 server.command = Some(pathext_command.to_string());
664 server.env.insert("PATHEXT".to_string(), ".CMD".to_string());
665 assert_eq!(
666 static_mcp_command_availability(&server).expect("PATHEXT command check"),
667 McpCommandAvailability::NotChecked,
668 "a child-PATH miss is conservative because Windows still searches implicit fallbacks"
669 );
670 server.command = Some(format!("{pathext_command}.cmd"));
671 assert_eq!(
672 static_mcp_command_availability(&server).expect("explicit .cmd command check"),
673 McpCommandAvailability::Available,
674 "Rust requires non-.exe extensions to be explicit"
675 );
676 }
677
678 #[tokio::test]
679 async fn mcp_http_auth_prefers_static_authorization_over_bearer_env() {
680 let mut headers = HashMap::new();
681 headers.insert("Authorization".to_string(), "Bearer static".to_string());
682 let auth = McpHttpAuth {
683 headers,
684 bearer_token_env_var: Some("PATH".to_string()),
685 ..Default::default()
686 };
687
688 let resolved = auth.resolved_headers().await.unwrap();
689 assert_eq!(
690 resolved.get("Authorization"),
691 Some(&"Bearer static".to_string())
692 );
693 }
694
695 #[tokio::test]
696 async fn mcp_http_auth_uses_bearer_env_when_no_authorization_header() {
697 let auth = McpHttpAuth {
698 bearer_token_env_var: Some("PATH".to_string()),
699 ..Default::default()
700 };
701
702 let resolved = auth.resolved_headers().await.unwrap();
703 assert!(
704 resolved
705 .get("Authorization")
706 .is_some_and(|value| value.starts_with("Bearer ") && value.len() > "Bearer ".len()),
707 "expected PATH-backed bearer header, got {resolved:?}"
708 );
709 }
710
711 #[test]
712 fn is_safe_custom_header_accepts_normal_auth_pairs() {
713 assert!(is_safe_custom_header("Authorization", "Bearer tok"));
714 assert!(is_safe_custom_header("X-Api-Key", "deadbeef"));
715 assert!(is_safe_custom_header("x-org", "anthropic"));
716 }
717
718 #[test]
719 fn is_safe_custom_header_rejects_empty_or_whitespace_key() {
720 assert!(!is_safe_custom_header("", "value"));
721 assert!(!is_safe_custom_header(" ", "value"));
722 }
723
724 #[test]
725 fn is_safe_custom_header_rejects_response_splitting_values() {
726 assert!(
727 !is_safe_custom_header("X-Foo", "abc\r\nSet-Cookie: evil=1"),
728 "CRLF in value must reject — response-splitting defense"
729 );
730 assert!(
731 !is_safe_custom_header("X-Foo", "abc\nbar"),
732 "bare LF in value must reject"
733 );
734 assert!(
735 !is_safe_custom_header("X-Foo", "abc\rbar"),
736 "bare CR in value must reject"
737 );
738 }
739
740 #[test]
741 fn is_safe_custom_header_rejects_protocol_framing_overrides() {
742 // The MCP Streamable HTTP transport relies on its own
743 // Accept / Content-Type values for protocol negotiation;
744 // a stray user override would silently break tool discovery.
745 assert!(!is_safe_custom_header("Accept", "text/plain"));
746 assert!(!is_safe_custom_header("accept", "text/plain"));
747 assert!(!is_safe_custom_header("Content-Type", "text/plain"));
748 assert!(!is_safe_custom_header("CONTENT-TYPE", "x/y"));
749 }
750
751 #[test]
752 fn default_mcp_http_get_accepts_json_and_event_stream() {
753 let client = test_http_client();
754 let request = with_default_mcp_http_headers(client.get("https://example.invalid/mcp"), false)
755 .build()
756 .unwrap();
757 assert_eq!(
758 request.headers().get(ACCEPT).and_then(|v| v.to_str().ok()),
759 Some(MCP_HTTP_ACCEPT)
760 );
761 assert!(
762 request.headers().get(CONTENT_TYPE).is_none(),
763 "SSE GET requests should not advertise a JSON request body"
764 );
765 }
766
767 #[test]
768 fn default_mcp_http_post_accepts_json_and_event_stream() {
769 let client = test_http_client();
770 let request = with_default_mcp_http_headers(client.post("https://example.invalid/mcp"), true)
771 .build()
772 .unwrap();
773 assert_eq!(
774 request.headers().get(ACCEPT).and_then(|v| v.to_str().ok()),
775 Some(MCP_HTTP_ACCEPT)
776 );
777 assert_eq!(
778 request
779 .headers()
780 .get(CONTENT_TYPE)
781 .and_then(|v| v.to_str().ok()),
782 Some("application/json")
783 );
784 }
785
786 #[tokio::test]
787 async fn streamable_http_transport_prepares_configured_headers() {
788 let mut headers = HashMap::new();
789 headers.insert("Authorization".to_string(), "Bearer xyz".to_string());
790 let url = "https://example.invalid/mcp";
791 let transport = StreamableHttpTransport::new(
792 test_mcp_http_client(url).with_mcp_auth(McpHttpAuth {
793 headers,
794 ..Default::default()
795 }),
796 url.to_string(),
797 );
798 let request = transport
799 .client
800 .prepare_mcp_request(transport.client.post(&transport.url), true)
801 .await
802 .unwrap()
803 .build()
804 .unwrap();
805 assert_eq!(
806 request.headers().get("Authorization").unwrap(),
807 "Bearer xyz"
808 );
809 assert_eq!(request.headers().get(ACCEPT).unwrap(), MCP_HTTP_ACCEPT);
810 assert_eq!(
811 request.headers().get(CONTENT_TYPE).unwrap(),
812 "application/json"
813 );
814 }
815
816 #[test]
817 fn mcp_auth_required_error_item_is_model_visible() {
818 let pool = McpPool::new(McpConfig::default());
819 let item = pool.mcp_auth_required_error_item("nordic-mcp");
820 assert_eq!(item["error"], "authentication_required");
821 assert_eq!(item["server"], "nordic-mcp");
822 assert!(
823 item.get("authenticate_tool").is_none(),
824 "an OAuth-servable synthetic tool is only named for a configured needs-auth server: {item}"
825 );
826 assert!(
827 item["message"]
828 .as_str()
829 .expect("message")
830 .contains("codewhale mcp login nordic-mcp")
831 );
832 }
833
834 #[test]
835 fn test_mcp_config_parse_mcp_servers_alias_and_snapshot() {
836 let dir = tempfile::tempdir().unwrap();
837 let path = dir.path().join("mcp.json");
838 fs::write(
839 &path,
840 r#"{
841 "mcpServers": {
842 "disabled": {
843 "command": "node",
844 "args": ["server.js"],
845 "disabled": true
846 }
847 }
848 }"#,
849 )
850 .unwrap();
851
852 let cfg = load_config(&path).unwrap();
853 assert!(cfg.servers.contains_key("disabled"));
854 let snapshot = manager_snapshot_from_config(&path, true).unwrap();
855 assert!(snapshot.reload_required);
856 assert_eq!(snapshot.servers[0].name, "disabled");
857 assert!(!snapshot.servers[0].enabled);
858 assert_eq!(snapshot.servers[0].error.as_deref(), Some("disabled"));
859 assert_eq!(
860 snapshot.servers[0].capability_metadata,
861 McpServerCapabilityMetadata::NotObserved
862 );
863 }
864
865 #[test]
866 fn malformed_mcp_config_error_omits_secret_contents_and_keys() {
867 let dir = tempfile::tempdir().unwrap();
868 let path = dir.path().join("mcp.json");
869 let secret = "cw-secret-mcp-config-4507";
870 fs::write(
871 &path,
872 format!(
873 r#"{{"servers":{{"private":{{"headers":{{"Authorization":"{secret}"}} trailing-junk}}}}}}"#
874 ),
875 )
876 .unwrap();
877
878 let error = load_config(&path).expect_err("malformed MCP config must fail");
879 let diagnostic = format!("{error:#}");
880 assert!(!diagnostic.contains(secret), "{diagnostic}");
881 assert!(!diagnostic.contains("Authorization"), "{diagnostic}");
882 assert!(
883 diagnostic.contains("file contents were omitted"),
884 "{diagnostic}"
885 );
886 }
887
888 #[test]
889 fn workspace_mcp_config_merges_with_project_overrides() {
890 let dir = tempfile::tempdir().unwrap();
891 let global_path = dir.path().join("global-mcp.json");
892 let workspace = dir.path().join("workspace");
893 let project_dir = workspace.join(".codewhale");
894 fs::create_dir_all(&project_dir).unwrap();
895 let _trust = mark_workspace_trusted(&workspace);
896 fs::write(
897 &global_path,
898 r#"{
899 "servers": {
900 "global": {"command": "node", "args": ["global.js"]},
901 "shared": {"command": "node", "args": ["global-shared.js"]}
902 }
903 }"#,
904 )
905 .unwrap();
906 fs::write(
907 project_dir.join("mcp.json"),
908 r#"{
909 "servers": {
910 "project": {"command": "php", "args": ["artisan", "boost:mcp"]},
911 "shared": {"command": "php", "args": ["artisan", "shared:mcp"]}
912 }
913 }"#,
914 )
915 .unwrap();
916
917 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
918 let workspace = workspace.canonicalize().unwrap();
919
920 assert!(cfg.servers.contains_key("global"));
921 let project = cfg.servers.get("project").unwrap();
922 assert_eq!(project.command.as_deref(), Some("php"));
923 assert_eq!(project.cwd.as_deref(), Some(workspace.as_path()));
924 let shared = cfg.servers.get("shared").unwrap();
925 assert_eq!(shared.args, vec!["artisan", "shared:mcp"]);
926 assert_eq!(shared.cwd.as_deref(), Some(workspace.as_path()));
927 }
928
929 #[test]
930 fn workspace_manager_snapshot_counts_global_and_project_servers() {
931 let dir = tempfile::tempdir().unwrap();
932 let global_path = dir.path().join("global-mcp.json");
933 let workspace = dir.path().join("workspace");
934 let project_dir = workspace.join(".codewhale");
935 fs::create_dir_all(&project_dir).unwrap();
936 let _trust = mark_workspace_trusted(&workspace);
937 fs::write(
938 &global_path,
939 r#"{
940 "servers": {
941 "chrome-devtools": {"command": "npx", "args": ["-y", "chrome-devtools-mcp@latest"]},
942 "context7": {"command": "npx", "args": ["-y", "@upstash/context7-mcp@latest"]}
943 }
944 }"#,
945 )
946 .unwrap();
947 fs::write(
948 project_dir.join("mcp.json"),
949 r#"{
950 "servers": {
951 "laravel-boost": {"command": "php", "args": ["artisan", "boost:mcp"]}
952 }
953 }"#,
954 )
955 .unwrap();
956
957 let plain = manager_snapshot_from_config(&global_path, false).unwrap();
958 let merged =
959 manager_snapshot_from_config_with_workspace(&global_path, &workspace, false).unwrap();
960
961 assert_eq!(plain.servers.len(), 2);
962 assert_eq!(merged.servers.len(), 3);
963 assert!(
964 merged
965 .servers
966 .iter()
967 .any(|server| server.name == "laravel-boost"),
968 "workspace-aware snapshots must include trusted project MCP servers"
969 );
970 }
971
972 #[test]
973 fn plugin_mcp_servers_are_qualified_and_resolve_relative_cwd() {
974 let dir = tempfile::tempdir().unwrap();
975 let plugin_base = dir.path().join("plugins").join("fleet");
976 fs::create_dir_all(plugin_base.join("servers/local")).unwrap();
977 fs::write(plugin_base.join("servers/local/server.js"), "// server\n").unwrap();
978
979 fs::write(
980 plugin_base.join("plugin.toml"),
981 r#"
982 schema_version = 1
983 [plugin]
984 name = "fleet"
985 version = "1.0.0"
986
987 [mcp_servers.local]
988 command = "node"
989 args = ["server.js"]
990 cwd = "servers/local"
991
992 [mcp_servers.remote]
993 url = "https://example.invalid/mcp"
994
995 [capabilities]
996 network_hosts = ["example.invalid"]
997 "#,
998 )
999 .unwrap();
1000 let (plugin, authority) = active_plugin_fixture(&plugin_base);
1001 let plugin_for_collision = plugin.clone();
1002 let authority_for_collision = authority.clone();
1003 let mut config = McpConfig::default();
1004 config.servers.insert(
1005 "global".to_string(),
1006 serde_json::from_str(r#"{"command":"node","args":["global.js"]}"#).unwrap(),
1007 );
1008
1009 let cfg = merge_plugin_mcp_servers_from_plugins(
1010 config,
1011 vec![("fleet".to_string(), plugin, authority)],
1012 )
1013 .unwrap();
1014
1015 assert!(cfg.servers.contains_key("global"));
1016
1017 let local = cfg.servers.get("plugin-5-fleet-local").unwrap();
1018 assert_eq!(local.command.as_deref(), Some("node"));
1019 let staged_root = plugin_for_collision.staged_root.as_deref().unwrap();
1020 assert_eq!(
1021 local.args,
1022 vec![
1023 staged_root
1024 .join("servers/local/server.js")
1025 .display()
1026 .to_string()
1027 ]
1028 );
1029 assert_eq!(
1030 local.cwd.as_deref(),
1031 Some(staged_root.join("servers/local").as_path())
1032 );
1033
1034 let remote = cfg.servers.get("plugin-5-fleet-remote").unwrap();
1035 assert_eq!(remote.url.as_deref(), Some("https://example.invalid/mcp"));
1036 assert!(remote.cwd.is_none());
1037
1038 let mut explicit = McpConfig::default();
1039 explicit.servers.insert(
1040 "plugin-5-fleet-local".to_string(),
1041 serde_json::from_str(r#"{"command":"node","args":["explicit.js"]}"#).unwrap(),
1042 );
1043 let collision_safe = merge_plugin_mcp_servers_from_plugins(
1044 explicit,
1045 vec![(
1046 "fleet".to_string(),
1047 plugin_for_collision,
1048 authority_for_collision,
1049 )],
1050 )
1051 .unwrap();
1052 assert_eq!(
1053 collision_safe.servers["plugin-5-fleet-local"].args,
1054 vec!["explicit.js"],
1055 "explicit MCP config must outrank a colliding plugin server"
1056 );
1057 }
1058
1059 #[cfg(target_os = "macos")]
1060 #[tokio::test]
1061 async fn reviewed_node_mjs_plugin_connects_through_inherited_descriptor() {
1062 if std::process::Command::new("node")
1063 .arg("--version")
1064 .output()
1065 .is_err()
1066 {
1067 eprintln!("skipping reviewed Node ESM launch test because node is unavailable");
1068 return;
1069 }
1070
1071 let dir = tempfile::tempdir().unwrap();
1072 let plugins_root = dir.path().join("plugins");
1073 let plugin_base = plugins_root.join("node-esm");
1074 fs::create_dir_all(&plugin_base).unwrap();
1075 fs::write(
1076 plugin_base.join("server.mjs"),
1077 r#"import readline from 'node:readline';
1078 const lines = readline.createInterface({ input: process.stdin });
1079 lines.on('line', (line) => {
1080 const request = JSON.parse(line);
1081 if (request.id === undefined) return;
1082 let result;
1083 if (request.method === 'initialize') {
1084 result = {
1085 protocolVersion: '2025-06-18',
1086 capabilities: { tools: {} },
1087 serverInfo: { name: 'node-esm', version: '1.0.0' }
1088 };
1089 } else if (request.method === 'tools/list') {
1090 result = {
1091 tools: [{ name: 'ready', description: 'ready', inputSchema: { type: 'object' } }]
1092 };
1093 } else {
1094 result = {};
1095 }
1096 process.stdout.write(JSON.stringify({ jsonrpc: '2.0', id: request.id, result }) + '\n');
1097 });
1098 "#,
1099 )
1100 .unwrap();
1101 fs::write(
1102 plugin_base.join("plugin.toml"),
1103 r#"
1104 schema_version = 1
1105 [plugin]
1106 name = "node-esm"
1107 version = "1.0.0"
1108
1109 [mcp_servers.local]
1110 command = "node"
1111 args = ["server.mjs"]
1112 connect_timeout = 2
1113 "#,
1114 )
1115 .unwrap();
1116
1117 let discovery = crate::plugins::discovery::DiscoveryConfig {
1118 workspace: dir.path().join("project"),
1119 user_plugins_dir: plugins_root,
1120 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
1121 builtin_plugin_dirs: Vec::new(),
1122 state_path: dir.path().join("plugin-state/state.json"),
1123 };
1124 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1125 registry.trust("node-esm").unwrap();
1126 registry.enable("node-esm").unwrap();
1127 let active = registry.active_plugins()[0].clone();
1128 let authority = registry.authority_for("node-esm").unwrap();
1129 let merged = merge_plugin_mcp_servers_from_plugins(
1130 McpConfig::default(),
1131 vec![("node-esm".to_string(), active, authority)],
1132 )
1133 .unwrap();
1134 let mut pool = McpPool::new(merged);
1135
1136 let connection = pool
1137 .get_or_connect("plugin-8-node-esm-local")
1138 .await
1139 .unwrap();
1140 assert_eq!(connection.tools().len(), 1);
1141 assert_eq!(connection.tools()[0].name, "ready");
1142 }
1143
1144 #[cfg(target_os = "macos")]
1145 #[tokio::test]
1146 async fn reviewed_node_plugins_preserve_module_path_context() {
1147 if std::process::Command::new("node")
1148 .arg("--version")
1149 .output()
1150 .is_err()
1151 {
1152 eprintln!("skipping reviewed multi-file Node ESM launch test because node is unavailable");
1153 return;
1154 }
1155
1156 for (extension, package_type) in [
1157 ("mjs", "module"),
1158 ("js", "module"),
1159 ("js", "commonjs"),
1160 ("cjs", "module"),
1161 ] {
1162 let esm = extension != "cjs" && package_type == "module";
1163 // The entry imports a sibling module, exactly like the computer-use
1164 // bundle (#5916). Launched by descriptor, Node would resolve `./lib/...`
1165 // against `/dev/` and the child would die before the handshake.
1166 let dir = tempfile::tempdir().unwrap();
1167 let plugins_root = dir.path().join("plugins");
1168 let plugin_base = plugins_root.join("node-esm-multi");
1169 fs::create_dir_all(plugin_base.join("mcp")).unwrap();
1170 fs::create_dir_all(plugin_base.join("lib")).unwrap();
1171 fs::write(
1172 plugin_base.join("package.json"),
1173 format!(r#"{{"type":"{package_type}"}}"#),
1174 )
1175 .unwrap();
1176 fs::write(plugin_base.join("mcp/reply.json"), r#"{"answer":42}"#).unwrap();
1177 fs::write(
1178 plugin_base.join("lib").join(format!("reply.{extension}")),
1179 if esm {
1180 r#"import path from 'node:path';
1181 import url from 'node:url';
1182 export const TOOL = 'ready-from-sibling';
1183 export const ENTRY_DIR = path.basename(path.dirname(url.fileURLToPath(import.meta.url)));
1184 "#
1185 } else {
1186 r#"const path = require('node:path');
1187 exports.TOOL = 'ready-from-sibling';
1188 exports.ENTRY_DIR = path.basename(__dirname);
1189 "#
1190 },
1191 )
1192 .unwrap();
1193 let imports = if esm {
1194 format!(
1195 "import readline from 'node:readline';\nimport fs from 'node:fs';\nimport {{ TOOL, ENTRY_DIR }} from '../lib/reply.{extension}';\n"
1196 )
1197 } else {
1198 format!(
1199 "const readline = require('node:readline');\nconst fs = require('node:fs');\nconst {{ TOOL, ENTRY_DIR }} = require('../lib/reply.{extension}');\n"
1200 )
1201 };
1202 fs::write(
1203 plugin_base.join("mcp").join(format!("server.{extension}")),
1204 imports
1205 + r#"const lines = readline.createInterface({ input: process.stdin });
1206 lines.on('line', (line) => {
1207 const request = JSON.parse(line);
1208 if (request.id === undefined) return;
1209 let result;
1210 if (request.method === 'initialize') {
1211 result = {
1212 protocolVersion: '2025-06-18',
1213 capabilities: { tools: {} },
1214 serverInfo: { name: 'node-esm-multi', version: '1.0.0' }
1215 };
1216 } else if (request.method === 'tools/list') {
1217 result = {
1218 tools: [{ name: TOOL, description: ENTRY_DIR, inputSchema: { type: 'object' } }]
1219 };
1220 } else if (request.method === 'tools/call') {
1221 result = { content: [{ type: 'text', text: JSON.stringify({
1222 answer: JSON.parse(fs.readFileSync('reply.json', 'utf8')).answer,
1223 args: process.argv.slice(2)
1224 }) }] };
1225 } else {
1226 result = {};
1227 }
1228 process.stdout.write(JSON.stringify({ jsonrpc: '2.0', id: request.id, result }) + '\n');
1229 });
1230 "#,
1231 )
1232 .unwrap();
1233 fs::write(
1234 plugin_base.join("plugin.toml"),
1235 format!(
1236 r#"
1237 schema_version = 1
1238 [plugin]
1239 name = "node-esm-multi"
1240 version = "1.0.0"
1241
1242 [mcp_servers.local]
1243 command = "node"
1244 args = ["--no-warnings", "server.{extension}", "fixture-argument"]
1245 cwd = "mcp"
1246 connect_timeout = 2
1247 "#
1248 ),
1249 )
1250 .unwrap();
1251
1252 let discovery = crate::plugins::discovery::DiscoveryConfig {
1253 workspace: dir.path().join("project"),
1254 user_plugins_dir: plugins_root,
1255 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
1256 builtin_plugin_dirs: Vec::new(),
1257 state_path: dir.path().join("plugin-state/state.json"),
1258 };
1259 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1260 registry.trust("node-esm-multi").unwrap();
1261 registry.enable("node-esm-multi").unwrap();
1262 let active = registry.active_plugins()[0].clone();
1263 let authority = registry.authority_for("node-esm-multi").unwrap();
1264 let merged = merge_plugin_mcp_servers_from_plugins(
1265 McpConfig::default(),
1266 vec![("node-esm-multi".to_string(), active, authority)],
1267 )
1268 .unwrap();
1269 let mut pool = McpPool::new(merged);
1270
1271 let connection = pool
1272 .get_or_connect("plugin-14-node-esm-multi-local")
1273 .await
1274 .unwrap();
1275 assert_eq!(connection.tools().len(), 1);
1276 assert_eq!(connection.tools()[0].name, "ready-from-sibling");
1277 // The sibling resolved from the staged tree, not from `/dev/`.
1278 assert_eq!(connection.tools()[0].description.as_deref(), Some("lib"));
1279 let result = pool
1280 .call_tool(
1281 "mcp_plugin-14-node-esm-multi-local_ready-from-sibling",
1282 serde_json::json!({}),
1283 )
1284 .await
1285 .unwrap();
1286 assert_eq!(
1287 result["content"][0]["text"], r#"{"answer":42,"args":["fixture-argument"]}"#,
1288 "{extension}/{package_type} must preserve staged cwd resources and script arguments"
1289 );
1290 registry.disable("node-esm-multi").unwrap();
1291 assert!(pool.all_tools().is_empty());
1292 assert!(
1293 pool.call_tool(
1294 "mcp_plugin-14-node-esm-multi-local_ready-from-sibling",
1295 serde_json::json!({})
1296 )
1297 .await
1298 .is_err(),
1299 "disabled Node tools must not remain callable"
1300 );
1301 }
1302 }
1303
1304 #[cfg(target_os = "macos")]
1305 #[test]
1306 fn node_entry_preserves_package_type_and_sibling_context() {
1307 use std::collections::BTreeMap;
1308 let staged_root = Path::new("/stage/plugin");
1309 let entry = staged_root.join("mcp/server.mjs");
1310 let hash = |paths: &[&str]| {
1311 paths
1312 .iter()
1313 .map(|path| (PathBuf::from(path), "h".to_string()))
1314 .collect::<BTreeMap<_, _>>()
1315 };
1316 // Even a single .js/.cjs depends on filename/package-type semantics.
1317 for extension in ["js", "cjs"] {
1318 let relative = format!("mcp/server.{extension}");
1319 assert!(node_entry_needs_staged_path(
1320 staged_root,
1321 &staged_root.join(&relative),
1322 &hash(&[&relative, "package.json"]),
1323 ));
1324 }
1325 // Manifests, docs, and data files are not modules.
1326 assert!(!node_entry_needs_staged_path(
1327 staged_root,
1328 &entry,
1329 &hash(&[
1330 "mcp/server.mjs",
1331 "plugin.json",
1332 "mcp.json",
1333 "README.md",
1334 "skills/a/SKILL.md"
1335 ]),
1336 ));
1337 for sibling in [
1338 "src/tools.mjs",
1339 "lib/x.js",
1340 "lib/x.cjs",
1341 "native/x.node",
1342 "wasm/x.wasm",
1343 ] {
1344 assert!(
1345 node_entry_needs_staged_path(
1346 staged_root,
1347 &entry,
1348 &hash(&["mcp/server.mjs", "plugin.json", sibling]),
1349 ),
1350 "{sibling} must force a path launch"
1351 );
1352 }
1353 // An entry outside the stage never qualifies.
1354 assert!(!node_entry_needs_staged_path(
1355 Path::new("/elsewhere"),
1356 &entry,
1357 &hash(&["mcp/server.mjs", "src/tools.mjs"]),
1358 ));
1359 }
1360
1361 #[cfg(target_os = "macos")]
1362 #[test]
1363 fn node_esm_descriptor_launch_keeps_options_argv_shape_and_script_arguments() {
1364 use std::ffi::OsString;
1365 let os = |value: &str| OsString::from(value);
1366
1367 // Options before the entrypoint stay in front; the entry is imported by
1368 // descriptor once, an empty main is supplied, and the descriptor path is
1369 // echoed after `--` so `process.argv[1]` keeps the file-mode shape.
1370 let args = vec![
1371 os("--max-old-space-size=256"),
1372 os("/dev/fd/7"),
1373 os("--port"),
1374 os("0"),
1375 ];
1376 assert_eq!(
1377 super::node_esm_descriptor_args(&args, 1),
1378 vec![
1379 os("--max-old-space-size=256"),
1380 os("--import"),
1381 os("/dev/fd/7"),
1382 os("-e"),
1383 os(""),
1384 os("--"),
1385 os("/dev/fd/7"),
1386 os("--port"),
1387 os("0"),
1388 ]
1389 );
1390
1391 // A `.mjs` that is not the first positional argument is not the
1392 // entrypoint; the launch is left untouched.
1393 let args = vec![os("other.js"), os("/dev/fd/7")];
1394 assert_eq!(super::node_esm_descriptor_args(&args, 1), args);
1395
1396 // The original option terminator cannot precede the injected --import.
1397 let args = vec![
1398 os("--no-warnings"),
1399 os("--"),
1400 os("/dev/fd/7"),
1401 os("argument"),
1402 ];
1403 assert_eq!(
1404 super::node_esm_descriptor_args(&args, 2),
1405 vec![
1406 os("--no-warnings"),
1407 os("--import"),
1408 os("/dev/fd/7"),
1409 os("-e"),
1410 os(""),
1411 os("--"),
1412 os("/dev/fd/7"),
1413 os("argument")
1414 ]
1415 );
1416 let args = vec![os("--conditions"), os("fixture"), os("/dev/fd/7")];
1417 let rewritten = super::node_esm_descriptor_args(&args, 2);
1418 assert_eq!(
1419 &rewritten[..4],
1420 &[
1421 os("--conditions"),
1422 os("fixture"),
1423 os("--import"),
1424 os("/dev/fd/7")
1425 ]
1426 );
1427 }
1428
1429 #[cfg(target_os = "macos")]
1430 #[test]
1431 fn reviewed_node_launch_identifies_only_the_script_operand() {
1432 for (args, expected) in [
1433 (vec!["/stage/server.js", "/stage/later.mjs"], Some(0)),
1434 (vec!["other.js", "/stage/later.mjs"], Some(0)),
1435 (
1436 vec!["--require", "/stage/preload.cjs", "/stage/server.js"],
1437 Some(2),
1438 ),
1439 (
1440 vec!["--import", "/stage/preload.mjs", "/stage/server.mjs"],
1441 Some(2),
1442 ),
1443 (vec!["--require"], None),
1444 (vec!["--", "/stage/server.cjs"], Some(1)),
1445 (vec!["--"], None),
1446 (
1447 vec!["--max-old-space-size=256", "/stage/server.mjs"],
1448 Some(1),
1449 ),
1450 (
1451 vec!["--max-old-space-size", "256", "/stage/server.mjs"],
1452 Some(2),
1453 ),
1454 (
1455 vec![
1456 "--abort-on-uncaught-exception",
1457 "--expose-gc",
1458 "--jitless",
1459 "/stage/server.mjs",
1460 ],
1461 Some(3),
1462 ),
1463 (vec!["-e", "console.log('x')", "/stage/later.mjs"], None),
1464 (vec!["--eval=console.log('x')", "/stage/later.mjs"], None),
1465 (vec!["--run=task", "--", "/stage/later.js"], None),
1466 (vec!["--input-type=module", "/stage/server.mjs"], None),
1467 (vec!["--input-type", "module", "/stage/server.mjs"], None),
1468 (vec!["--unknown-option", "/stage/value.js"], None),
1469 (vec!["-", "/stage/later.mjs"], None),
1470 ] {
1471 assert_eq!(node_script_entry_index(&args), expected, "{args:?}");
1472 }
1473 }
1474
1475 #[test]
1476 fn plugin_server_ids_are_unambiguous_across_hyphenated_plugin_and_server_names() {
1477 let left = qualified_plugin_server_name("foo-bar", "baz");
1478 let right = qualified_plugin_server_name("foo", "bar-baz");
1479
1480 assert_eq!(left, "plugin-7-foo-bar-baz");
1481 assert_eq!(right, "plugin-3-foo-bar-baz");
1482 assert_ne!(left, right);
1483 }
1484
1485 #[test]
1486 fn mixed_components_contribute_only_mcp_servers_to_the_mcp_catalog() {
1487 let dir = tempfile::tempdir().unwrap();
1488 let plugin_base = dir.path().join("plugin");
1489 fs::create_dir_all(plugin_base.join("commands")).unwrap();
1490 fs::create_dir_all(plugin_base.join("hooks")).unwrap();
1491 fs::write(plugin_base.join("server.js"), "// reviewed entrypoint\n").unwrap();
1492 fs::write(
1493 plugin_base.join("plugin.toml"),
1494 r#"
1495 schema_version = 1
1496 [plugin]
1497 name = "fleet"
1498 version = "1.0.0"
1499
1500 [mcp_servers.local]
1501 command = "node"
1502 args = ["server.js"]
1503
1504 [mcp_servers.remote]
1505 url = "https://example.invalid/mcp"
1506
1507 [capabilities]
1508 network_hosts = ["example.invalid"]
1509
1510 [commands]
1511 path = "commands"
1512
1513 [hooks]
1514 path = "hooks"
1515 "#,
1516 )
1517 .unwrap();
1518 let (plugin, authority) = active_plugin_fixture(&plugin_base);
1519 assert!(plugin.active());
1520
1521 let cfg = merge_plugin_mcp_servers_from_plugins(
1522 McpConfig::default(),
1523 vec![("fleet".to_string(), plugin.clone(), authority)],
1524 )
1525 .unwrap();
1526 assert!(cfg.servers.contains_key("plugin-5-fleet-local"));
1527 assert!(cfg.servers.contains_key("plugin-5-fleet-remote"));
1528 assert_eq!(
1529 cfg.servers.len(),
1530 2,
1531 "declarative components must not become MCP servers: {:?}",
1532 cfg.servers.keys().collect::<Vec<_>>()
1533 );
1534 assert_eq!(
1535 cfg.servers["plugin-5-fleet-local"].command.as_deref(),
1536 Some("node")
1537 );
1538 assert_eq!(
1539 cfg.servers["plugin-5-fleet-remote"].url.as_deref(),
1540 Some("https://example.invalid/mcp")
1541 );
1542 }
1543
1544 #[test]
1545 fn plugin_mcp_adapter_denies_disabled_and_untrusted_bundles() {
1546 let dir = tempfile::tempdir().unwrap();
1547 let plugin_base = dir.path().join("plugin");
1548 fs::create_dir_all(&plugin_base).unwrap();
1549 fs::write(
1550 plugin_base.join("plugin.toml"),
1551 r#"
1552 schema_version = 1
1553 [plugin]
1554 name = "denied"
1555 version = "1.0.0"
1556
1557 [mcp_servers.local]
1558 command = "node"
1559 "#,
1560 )
1561 .unwrap();
1562 let (mut disabled, authority) = active_plugin_fixture(&plugin_base);
1563 disabled.enabled = false;
1564 let mut untrusted = disabled.clone();
1565 untrusted.enabled = true;
1566 untrusted.trust_status = crate::plugins::types::PluginTrustStatus::NeverReviewed;
1567
1568 for plugin in [disabled, untrusted] {
1569 let config = merge_plugin_mcp_servers_from_plugins(
1570 McpConfig::default(),
1571 vec![("denied".to_string(), plugin, authority.clone())],
1572 )
1573 .unwrap();
1574 assert!(
1575 config.servers.is_empty(),
1576 "headless MCP adapter admitted an inactive bundle"
1577 );
1578 }
1579 }
1580
1581 #[test]
1582 fn plugin_mcp_adapter_denies_content_changed_after_snapshot() {
1583 let dir = tempfile::tempdir().unwrap();
1584 let plugin_base = dir.path().join("plugin");
1585 fs::create_dir_all(&plugin_base).unwrap();
1586 let manifest_path = plugin_base.join("plugin.toml");
1587 fs::write(
1588 &manifest_path,
1589 r#"
1590 schema_version = 1
1591 [plugin]
1592 name = "changed"
1593 version = "1.0.0"
1594
1595 [mcp_servers.local]
1596 command = "node"
1597 "#,
1598 )
1599 .unwrap();
1600 let (plugin, authority) = active_plugin_fixture(&plugin_base);
1601 fs::write(plugin_base.join("late-change.txt"), "changed after review").unwrap();
1602
1603 let config = merge_plugin_mcp_servers_from_plugins(
1604 McpConfig::default(),
1605 vec![("changed".to_string(), plugin, authority)],
1606 )
1607 .unwrap();
1608 assert!(config.servers.is_empty());
1609 }
1610
1611 fn registry_with_local_mcp(
1612 name: &str,
1613 base_path: PathBuf,
1614 workspace: &Path,
1615 ) -> crate::plugins::PluginRegistry {
1616 fs::write(
1617 base_path.join("plugin.toml"),
1618 format!(
1619 r#"
1620 schema_version = 1
1621 [plugin]
1622 name = "{name}"
1623 version = "1.0.0"
1624
1625 [mcp_servers.local]
1626 command = "node"
1627 args = ["server.js"]
1628 "#,
1629 ),
1630 )
1631 .unwrap();
1632 let plugins_root = base_path.parent().expect("plugin parent").to_path_buf();
1633 let discovery = crate::plugins::discovery::DiscoveryConfig {
1634 workspace: workspace.to_path_buf(),
1635 user_plugins_dir: plugins_root,
1636 workspace_plugins_dir: workspace.join(".codewhale/plugins-unused"),
1637 builtin_plugin_dirs: Vec::new(),
1638 state_path: workspace
1639 .join("plugin-state")
1640 .join(format!("plugin-state-{name}.json")),
1641 };
1642 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1643 registry.trust(name).unwrap();
1644 registry.enable(name).unwrap();
1645 registry
1646 }
1647
1648 #[test]
1649 fn plugin_mcp_servers_merge_without_project_config() {
1650 let dir = tempfile::tempdir().unwrap();
1651 let global_path = dir.path().join("global-mcp.json");
1652 let workspace = dir.path().join("workspace");
1653 let plugin_base = dir.path().join("plugins").join("fixture");
1654 fs::create_dir_all(&workspace).unwrap();
1655 fs::create_dir_all(&plugin_base).unwrap();
1656 fs::write(
1657 &global_path,
1658 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
1659 )
1660 .unwrap();
1661
1662 let plugins = registry_with_local_mcp("fixture", plugin_base.clone(), &workspace);
1663 let staged_root = plugins
1664 .get("fixture")
1665 .and_then(|plugin| plugin.staged_root.clone())
1666 .expect("trusted plugin should have an immutable runtime snapshot");
1667 let cfg = load_config_with_workspace_and_plugins(&global_path, &workspace, &plugins).unwrap();
1668
1669 assert!(cfg.servers.contains_key("global"));
1670 let qualified_name = qualified_plugin_server_name("fixture", "local");
1671 let local = cfg
1672 .servers
1673 .get(&qualified_name)
1674 .expect("plugin MCP should merge without a project MCP config");
1675 assert_eq!(local.command.as_deref(), Some("node"));
1676 assert_eq!(local.cwd.as_deref(), Some(staged_root.as_path()));
1677 }
1678
1679 #[cfg(unix)]
1680 #[tokio::test]
1681 async fn plugin_mcp_lazy_spawn_denies_component_changed_after_pool_construction() {
1682 use std::os::unix::fs::PermissionsExt;
1683
1684 let dir = tempfile::tempdir().unwrap();
1685 let plugins_root = dir.path().join("plugins");
1686 let plugin_base = plugins_root.join("guarded");
1687 fs::create_dir_all(&plugin_base).unwrap();
1688 let server_path = plugin_base.join("server.sh");
1689 fs::write(&server_path, "#!/bin/sh\nexit 0\n").unwrap();
1690 let mut permissions = fs::metadata(&server_path).unwrap().permissions();
1691 permissions.set_mode(0o700);
1692 fs::set_permissions(&server_path, permissions).unwrap();
1693 fs::write(
1694 plugin_base.join("plugin.toml"),
1695 r#"
1696 schema_version = 1
1697 [plugin]
1698 name = "guarded"
1699 version = "1.0.0"
1700
1701 [mcp_servers.local]
1702 command = "sh"
1703 args = ["server.sh"]
1704 connect_timeout = 1
1705 "#,
1706 )
1707 .unwrap();
1708
1709 let discovery = crate::plugins::discovery::DiscoveryConfig {
1710 workspace: dir.path().join("project"),
1711 user_plugins_dir: plugins_root,
1712 workspace_plugins_dir: dir.path().join("workspace-plugins"),
1713 builtin_plugin_dirs: Vec::new(),
1714 state_path: dir.path().join("plugin-state/state.json"),
1715 };
1716 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1717 registry.trust("guarded").unwrap();
1718 registry.enable("guarded").unwrap();
1719 let active = registry.active_plugins()[0].clone();
1720 let authority = registry.authority_for("guarded").unwrap();
1721 let merged = merge_plugin_mcp_servers_from_plugins(
1722 McpConfig::default(),
1723 vec![("guarded".to_string(), active, authority)],
1724 )
1725 .unwrap();
1726 assert!(
1727 merged.servers["plugin-7-guarded-local"]
1728 .reviewed_plugin
1729 .is_some(),
1730 "plugin provenance must survive through MCP pool construction"
1731 );
1732 let mut pool = McpPool::new(merged);
1733
1734 // Adversarial mutation after trust, enablement, merge, and pool
1735 // construction. If the lazy child executes, it creates this marker before
1736 // closing stdio, so the regression proves denial happened pre-spawn.
1737 let executed_marker = plugin_base.join("executed.marker");
1738 fs::write(&server_path, "#!/bin/sh\n: > executed.marker\nexit 0\n").unwrap();
1739
1740 let error = pool
1741 .get_or_connect("plugin-7-guarded-local")
1742 .await
1743 .err()
1744 .expect("changed reviewed component must be denied before spawn");
1745 let message = format!("{error:#}");
1746 assert!(
1747 message.contains("Refusing to use MCP server 'plugin-7-guarded-local'"),
1748 "unexpected pre-spawn denial: {message}"
1749 );
1750 assert!(message.contains("changed after review"));
1751 assert!(message.contains("/plugin reload"));
1752 assert!(
1753 !executed_marker.exists(),
1754 "mutated MCP component executed despite pre-spawn hash denial"
1755 );
1756 }
1757
1758 #[cfg(unix)]
1759 #[tokio::test]
1760 async fn plugin_mcp_inflight_call_is_cancelled_after_cross_process_revocation() {
1761 let _env_lock = crate::test_support::lock_test_env();
1762 let dir = tempfile::tempdir().unwrap();
1763 let call_marker = dir.path().join("call.marker");
1764 let _call_marker_env = crate::test_support::EnvVarGuard::set(
1765 "CODEWHALE_TEST_PLUGIN_CALL_MARKER",
1766 call_marker.as_os_str(),
1767 );
1768 let plugins_root = dir.path().join("plugins");
1769 let plugin_base = plugins_root.join("revoked");
1770 fs::create_dir_all(&plugin_base).unwrap();
1771 fs::create_dir_all(dir.path().join("project")).unwrap();
1772 fs::write(
1773 plugin_base.join("server.sh"),
1774 r#"#!/bin/sh
1775 trap 'exit 0' TERM INT
1776 while IFS= read -r line; do
1777 case "$line" in
1778 *'"method":"notifications/initialized"'*)
1779 ;;
1780 *'"method":"initialize"'*)
1781 printf '%s\n' '{"jsonrpc":"2.0","id":"1","result":{"protocolVersion":"2024-11-05","serverInfo":{"name":"revocation-test","version":"1.0.0"},"capabilities":{"tools":{}}}}'
1782 ;;
1783 *'"method":"tools/list"'*)
1784 printf '%s\n' '{"jsonrpc":"2.0","id":"2","result":{"tools":[{"name":"wait","description":"Wait until revoked","inputSchema":{"type":"object"}}]}}'
1785 ;;
1786 *'"method":"tools/call"'*)
1787 : > "$CALL_MARKER"
1788 while :; do sleep 1; done
1789 ;;
1790 esac
1791 done
1792 "#,
1793 )
1794 .unwrap();
1795 fs::write(
1796 plugin_base.join("plugin.toml"),
1797 r#"
1798 schema_version = 1
1799 [plugin]
1800 name = "revoked"
1801 version = "1.0.0"
1802
1803 [mcp_servers.local]
1804 command = "sh"
1805 args = ["server.sh"]
1806 connect_timeout = 2
1807 execute_timeout = 30
1808 read_timeout = 30
1809
1810 [mcp_servers.local.env]
1811 CALL_MARKER = "${CODEWHALE_TEST_PLUGIN_CALL_MARKER}"
1812 "#,
1813 )
1814 .unwrap();
1815
1816 let discovery = crate::plugins::discovery::DiscoveryConfig {
1817 workspace: dir.path().join("project"),
1818 user_plugins_dir: plugins_root,
1819 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
1820 builtin_plugin_dirs: Vec::new(),
1821 state_path: dir.path().join("plugin-state/state.json"),
1822 };
1823 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
1824 registry.trust("revoked").unwrap();
1825 registry.enable("revoked").unwrap();
1826 let active = registry.active_plugins()[0].clone();
1827 let authority = registry.authority_for("revoked").unwrap();
1828 let merged = merge_plugin_mcp_servers_from_plugins(
1829 McpConfig::default(),
1830 vec![("revoked".to_string(), active, authority)],
1831 )
1832 .unwrap();
1833 let mut pool = McpPool::new(merged);
1834 pool.get_or_connect("plugin-7-revoked-local").await.unwrap();
1835
1836 let call = tokio::spawn(async move {
1837 pool.call_tool("mcp_plugin-7-revoked-local_wait", serde_json::json!({}))
1838 .await
1839 });
1840 for _ in 0..100 {
1841 if call_marker.exists() {
1842 break;
1843 }
1844 if call.is_finished() {
1845 let early = call
1846 .await
1847 .expect("in-flight tool task panicked before reaching the server");
1848 panic!("in-flight tool call ended before reaching the server: {early:?}");
1849 }
1850 tokio::time::sleep(Duration::from_millis(20)).await;
1851 }
1852 assert!(
1853 call_marker.exists(),
1854 "test server never observed the in-flight tool call"
1855 );
1856
1857 let mut external = crate::plugins::discovery::discover_with_config(&discovery);
1858 external.revoke_trust("revoked").unwrap();
1859 let result = tokio::time::timeout(Duration::from_secs(5), call)
1860 .await
1861 .expect("revocation watcher did not cancel the in-flight call")
1862 .unwrap();
1863 let error = result
1864 .expect_err("revoked in-flight call must not complete")
1865 .to_string();
1866 assert!(error.contains("cancelled after authority changed"));
1867 assert!(error.contains("disabled, revoked, or no longer matches"));
1868 }
1869
1870 #[cfg(unix)]
1871 #[tokio::test]
1872 async fn plugin_stdio_authority_cancellation_terminates_an_idle_child() {
1873 let dir = tempfile::tempdir().unwrap();
1874 let plugin_base = dir.path().join("plugins/idle-child");
1875 fs::create_dir_all(&plugin_base).unwrap();
1876 fs::write(
1877 plugin_base.join("plugin.toml"),
1878 "schema_version = 1\n[plugin]\nname = \"idle-child\"\nversion = \"1.0.0\"\n",
1879 )
1880 .unwrap();
1881 let (_, authority) = active_plugin_fixture(&plugin_base);
1882 let mut config = test_server_config();
1883 config.command = Some("sh".to_string());
1884 config.args = vec![
1885 "-c".to_string(),
1886 "trap 'exit 0' TERM INT; while :; do sleep 1; done".to_string(),
1887 ];
1888 config.reviewed_plugin = Some(
1889 ReviewedPluginMcpSource::from_authority(
1890 authority,
1891 None,
1892 Arc::new(crate::plugins::HostEnvironment::capture()),
1893 )
1894 .unwrap(),
1895 );
1896 let cancellation = tokio_util::sync::CancellationToken::new();
1897 let transport = StdioTransport::spawn(
1898 "idle-child",
1899 config.command.as_deref().unwrap(),
1900 &config,
1901 cancellation.clone(),
1902 )
1903 .unwrap();
1904 let child = transport.session.child_for_tests();
1905 assert!(child.lock().await.try_wait().unwrap().is_none());
1906
1907 cancellation.cancel();
1908 let deadline = tokio::time::Instant::now() + Duration::from_secs(3);
1909 loop {
1910 if child.lock().await.try_wait().unwrap().is_some() {
1911 break;
1912 }
1913 assert!(
1914 tokio::time::Instant::now() < deadline,
1915 "authority cancellation left the plugin stdio child alive"
1916 );
1917 tokio::time::sleep(Duration::from_millis(20)).await;
1918 }
1919 }
1920
1921 #[cfg(unix)]
1922 #[tokio::test]
1923 async fn plugin_stdio_does_not_surface_reviewed_child_stderr() {
1924 let dir = tempfile::tempdir().unwrap();
1925 let plugin_base = dir.path().join("plugins/stderr-secret");
1926 fs::create_dir_all(&plugin_base).unwrap();
1927 fs::write(
1928 plugin_base.join("plugin.toml"),
1929 "schema_version = 1\n[plugin]\nname = \"stderr-secret\"\nversion = \"1.0.0\"\n",
1930 )
1931 .unwrap();
1932 let (_, authority) = active_plugin_fixture(&plugin_base);
1933 let mut config = test_server_config();
1934 config.command = Some("sh".to_string());
1935 config.args = vec![
1936 "-c".to_string(),
1937 "echo 'ARBITRARY_PLUGIN_CREDENTIAL' 1>&2; exit 1".to_string(),
1938 ];
1939 config.reviewed_plugin = Some(
1940 ReviewedPluginMcpSource::from_authority(
1941 authority,
1942 None,
1943 Arc::new(crate::plugins::HostEnvironment::capture()),
1944 )
1945 .unwrap(),
1946 );
1947 let mut transport = StdioTransport::spawn(
1948 "stderr-secret",
1949 config.command.as_deref().unwrap(),
1950 &config,
1951 tokio_util::sync::CancellationToken::new(),
1952 )
1953 .unwrap();
1954
1955 tokio::time::sleep(Duration::from_millis(100)).await;
1956 let error = transport
1957 .recv()
1958 .await
1959 .expect_err("reviewed child should have closed its transport")
1960 .to_string();
1961 assert!(error.contains("Stdio transport closed"));
1962 assert!(!error.contains("ARBITRARY_PLUGIN_CREDENTIAL"));
1963 }
1964
1965 /// Shutdown reaches what the server started, not only the server: a
1966 /// background grandchild (an `npx` wrapper's node, a shell job) must not
1967 /// survive the transport.
1968 #[cfg(unix)]
1969 #[tokio::test]
1970 async fn stdio_shutdown_also_terminates_the_servers_grandchildren() {
1971 let mut config = test_server_config();
1972 config.command = Some("sh".to_string());
1973 config.args = vec!["-c".to_string(), "sleep 300 & echo $!; wait".to_string()];
1974 let mut transport = StdioTransport::spawn(
1975 "grandparent",
1976 "sh",
1977 &config,
1978 tokio_util::sync::CancellationToken::new(),
1979 )
1980 .unwrap();
1981 let line = tokio::time::timeout(Duration::from_secs(5), transport.recv())
1982 .await
1983 .expect("grandchild pid line")
1984 .unwrap();
1985 let grandchild: i32 = String::from_utf8(line).unwrap().trim().parse().unwrap();
1986
1987 transport.shutdown().await;
1988
1989 let deadline = tokio::time::Instant::now() + Duration::from_secs(3);
1990 loop {
1991 // SAFETY: signal 0 only probes whether the pid still exists.
1992 if unsafe { libc::kill(grandchild, 0) } != 0 {
1993 break;
1994 }
1995 if tokio::time::Instant::now() >= deadline {
1996 // Do not leak the sleeper past a failing run.
1997 // SAFETY: plain kill(2) of the pid this test started.
1998 unsafe {
1999 libc::kill(grandchild, libc::SIGKILL);
2000 }
2001 panic!("the MCP server's grandchild outlived the transport");
2002 }
2003 tokio::time::sleep(Duration::from_millis(20)).await;
2004 }
2005 }
2006
2007 /// #6187: a crashed stdio child must stop reading as "ready" before any
2008 /// call is in flight — `is_ready` probes the child, so the pool rebuilds
2009 /// the connection on the next use instead of handing the dead transport
2010 /// back.
2011 #[cfg(unix)]
2012 #[tokio::test]
2013 async fn dead_stdio_child_stops_reading_ready_without_a_call_in_flight() {
2014 let mut config = test_server_config();
2015 config.command = Some("sh".to_string());
2016 config.args = vec!["-c".to_string(), "while :; do sleep 1; done".to_string()];
2017 let transport = StdioTransport::spawn(
2018 "idle",
2019 "sh",
2020 &config,
2021 tokio_util::sync::CancellationToken::new(),
2022 )
2023 .unwrap();
2024 let child = transport.session.child_for_tests();
2025 let connection = test_connection(Box::new(transport));
2026
2027 // Alive child: the Ready state flag is the whole answer.
2028 assert!(
2029 connection.is_ready(),
2030 "a live stdio child must not be probed dead"
2031 );
2032
2033 child.lock().await.start_kill().unwrap();
2034 let deadline = tokio::time::Instant::now() + Duration::from_secs(3);
2035 loop {
2036 if child.lock().await.try_wait().unwrap().is_some() {
2037 break;
2038 }
2039 assert!(
2040 tokio::time::Instant::now() < deadline,
2041 "killed stdio child was never reaped"
2042 );
2043 tokio::time::sleep(Duration::from_millis(10)).await;
2044 }
2045
2046 assert!(
2047 !connection.is_ready(),
2048 "a reaped stdio child must fail is_ready without a call in flight"
2049 );
2050 }
2051
2052 #[tokio::test]
2053 async fn revoked_plugin_mcp_denies_catalog_tool_resource_and_prompt_operations() {
2054 let dir = tempfile::tempdir().unwrap();
2055 let plugins_root = dir.path().join("plugins");
2056 let plugin_base = plugins_root.join("catalog-guard");
2057 fs::create_dir_all(&plugin_base).unwrap();
2058 fs::create_dir_all(dir.path().join("project")).unwrap();
2059 fs::write(
2060 plugin_base.join("plugin.toml"),
2061 "schema_version = 1\n[plugin]\nname = \"catalog-guard\"\nversion = \"1.0.0\"\n",
2062 )
2063 .unwrap();
2064 let discovery = crate::plugins::discovery::DiscoveryConfig {
2065 workspace: dir.path().join("project"),
2066 user_plugins_dir: plugins_root,
2067 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
2068 builtin_plugin_dirs: Vec::new(),
2069 state_path: dir.path().join("plugin-state/state.json"),
2070 };
2071 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
2072 registry.trust("catalog-guard").unwrap();
2073 registry.enable("catalog-guard").unwrap();
2074 let authority = registry.authority_for("catalog-guard").unwrap();
2075
2076 let sent = Arc::new(Mutex::new(Vec::new()));
2077 let mut connection = test_connection(Box::new(ScriptedValueTransport {
2078 sent: Arc::clone(&sent),
2079 responses: VecDeque::new(),
2080 }));
2081 let source = ReviewedPluginMcpSource::from_authority(
2082 authority,
2083 None,
2084 Arc::new(crate::plugins::HostEnvironment::capture()),
2085 )
2086 .unwrap();
2087 connection.config.reviewed_plugin = Some(source.clone());
2088 connection.tools.push(McpTool {
2089 name: "echo".to_string(),
2090 description: None,
2091 input_schema: serde_json::json!({}),
2092 annotations: None,
2093 });
2094 connection.resources.push(McpResource {
2095 uri: "memory://one".to_string(),
2096 name: "one".to_string(),
2097 description: None,
2098 mime_type: None,
2099 });
2100 connection.resource_templates.push(McpResourceTemplate {
2101 uri_template: "memory://{id}".to_string(),
2102 name: "memory".to_string(),
2103 description: None,
2104 mime_type: None,
2105 });
2106 connection.prompts.push(McpPrompt {
2107 name: "review".to_string(),
2108 description: None,
2109 arguments: Vec::new(),
2110 });
2111 let mut config = McpConfig::default();
2112 let mut server = test_server_config();
2113 server.reviewed_plugin = Some(source);
2114 config.servers.insert("guarded".to_string(), server);
2115 let mut pool = McpPool::new(config);
2116 pool.connections.insert("guarded".to_string(), connection);
2117 assert_eq!(pool.all_tools().len(), 1);
2118 assert_eq!(pool.all_resources().len(), 1);
2119 assert_eq!(pool.all_resource_templates().len(), 1);
2120 assert_eq!(pool.all_prompts().len(), 1);
2121
2122 let mut external = crate::plugins::discovery::discover_with_config(&discovery);
2123 external.revoke_trust("catalog-guard").unwrap();
2124 assert!(pool.all_tools().is_empty());
2125 assert!(pool.all_resources().is_empty());
2126 assert!(pool.all_resource_templates().is_empty());
2127 assert!(pool.all_prompts().is_empty());
2128
2129 let tool = pool
2130 .call_tool("mcp_guarded_echo", serde_json::json!({}))
2131 .await;
2132 let resource = pool.read_resource("guarded", "memory://one").await;
2133 let prompt = pool
2134 .get_prompt("guarded", "review", serde_json::json!({}))
2135 .await;
2136 let resource_catalog = pool
2137 .call_tool(
2138 "list_mcp_resources",
2139 serde_json::json!({"server": "guarded"}),
2140 )
2141 .await;
2142 let template_catalog = pool
2143 .call_tool(
2144 "list_mcp_resource_templates",
2145 serde_json::json!({"server": "guarded"}),
2146 )
2147 .await;
2148 for result in [tool, resource, prompt, resource_catalog, template_catalog] {
2149 let error = result
2150 .expect_err("revoked plugin MCP operation must fail closed")
2151 .to_string();
2152 assert!(error.contains("Refusing to use MCP server 'guarded'"));
2153 }
2154 assert!(
2155 sent.lock().unwrap().is_empty(),
2156 "revoked plugin MCP operation reached the transport"
2157 );
2158 }
2159
2160 fn cached_reviewed_plugin_catalog_fixture() -> (tempfile::TempDir, PathBuf, PathBuf, McpPool) {
2161 let dir = tempfile::tempdir().unwrap();
2162 let plugins_root = dir.path().join("plugins");
2163 let plugin_base = plugins_root.join("catalog-drift");
2164 fs::create_dir_all(&plugin_base).unwrap();
2165 fs::create_dir_all(dir.path().join("project")).unwrap();
2166 fs::write(
2167 plugin_base.join("plugin.toml"),
2168 "schema_version = 1\n[plugin]\nname = \"catalog-drift\"\nversion = \"1.0.0\"\n",
2169 )
2170 .unwrap();
2171 let discovery = crate::plugins::discovery::DiscoveryConfig {
2172 workspace: dir.path().join("project"),
2173 user_plugins_dir: plugins_root,
2174 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
2175 builtin_plugin_dirs: Vec::new(),
2176 state_path: dir.path().join("plugin-state/state.json"),
2177 };
2178 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
2179 registry.trust("catalog-drift").unwrap();
2180 registry.enable("catalog-drift").unwrap();
2181 let authority = registry.authority_for("catalog-drift").unwrap();
2182 let staged_manifest = authority.staged_manifest.clone();
2183
2184 let mut connection = test_connection(Box::new(ScriptedValueTransport {
2185 sent: Arc::new(Mutex::new(Vec::new())),
2186 responses: VecDeque::new(),
2187 }));
2188 let source = ReviewedPluginMcpSource::from_authority(
2189 authority,
2190 None,
2191 Arc::new(crate::plugins::HostEnvironment::capture()),
2192 )
2193 .unwrap();
2194 connection.config.reviewed_plugin = Some(source.clone());
2195 connection.tools.push(McpTool {
2196 name: "echo".to_string(),
2197 description: None,
2198 input_schema: serde_json::json!({}),
2199 annotations: None,
2200 });
2201 connection.resources.push(McpResource {
2202 uri: "memory://one".to_string(),
2203 name: "one".to_string(),
2204 description: None,
2205 mime_type: None,
2206 });
2207 connection.resource_templates.push(McpResourceTemplate {
2208 uri_template: "memory://{id}".to_string(),
2209 name: "memory".to_string(),
2210 description: None,
2211 mime_type: None,
2212 });
2213 connection.prompts.push(McpPrompt {
2214 name: "review".to_string(),
2215 description: None,
2216 arguments: Vec::new(),
2217 });
2218 let mut config = McpConfig::default();
2219 let mut server = test_server_config();
2220 server.reviewed_plugin = Some(source);
2221 config.servers.insert("guarded".to_string(), server);
2222 let mut pool = McpPool::new(config);
2223 pool.connections.insert("guarded".to_string(), connection);
2224 assert_eq!(pool.all_tools().len(), 1);
2225 assert_eq!(pool.all_resources().len(), 1);
2226 assert_eq!(pool.all_resource_templates().len(), 1);
2227 assert_eq!(pool.all_prompts().len(), 1);
2228
2229 (dir, plugin_base, staged_manifest, pool)
2230 }
2231
2232 fn assert_reviewed_plugin_catalog_hidden(pool: &McpPool, boundary: &str) {
2233 assert!(pool.all_tools().is_empty());
2234 assert!(pool.all_resources().is_empty());
2235 assert!(pool.all_resource_templates().is_empty());
2236 assert!(pool.all_prompts().is_empty());
2237 assert!(
2238 pool.to_api_tools()
2239 .iter()
2240 .all(|tool| tool.name != "mcp_guarded_echo"),
2241 "{boundary} drift must remove cached reviewed tools from the model API catalog"
2242 );
2243 assert!(pool.parse_prefixed_name("mcp_guarded_echo").is_err());
2244 }
2245
2246 #[test]
2247 fn reviewed_plugin_source_drift_hides_every_cached_catalog_surface() {
2248 let (_dir, plugin_base, _staged_manifest, pool) = cached_reviewed_plugin_catalog_fixture();
2249
2250 fs::write(plugin_base.join("unreviewed-companion.txt"), b"drift").unwrap();
2251
2252 assert_reviewed_plugin_catalog_hidden(&pool, "source");
2253 }
2254
2255 #[cfg(unix)]
2256 #[test]
2257 fn reviewed_plugin_stage_drift_hides_every_cached_catalog_surface() {
2258 use std::io::Write as _;
2259 use std::os::unix::fs::PermissionsExt as _;
2260
2261 let (_dir, _plugin_base, staged_manifest, pool) = cached_reviewed_plugin_catalog_fixture();
2262 std::fs::set_permissions(&staged_manifest, std::fs::Permissions::from_mode(0o600)).unwrap();
2263 std::fs::OpenOptions::new()
2264 .append(true)
2265 .open(&staged_manifest)
2266 .unwrap()
2267 .write_all(b"\n# test-only staged drift\n")
2268 .unwrap();
2269
2270 assert_reviewed_plugin_catalog_hidden(&pool, "staged-tree");
2271 }
2272
2273 #[tokio::test]
2274 async fn reviewed_plugin_oauth_is_disabled_without_network_or_token_mutation() {
2275 let dir = tempfile::tempdir().unwrap();
2276 let plugin_base = dir.path().join("plugins/oauth-disabled");
2277 fs::create_dir_all(&plugin_base).unwrap();
2278 fs::write(
2279 plugin_base.join("plugin.toml"),
2280 "schema_version = 1\n[plugin]\nname = \"oauth-disabled\"\nversion = \"1.0.0\"\n",
2281 )
2282 .unwrap();
2283 let (_, authority) = active_plugin_fixture(&plugin_base);
2284
2285 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
2286 let endpoint = format!("http://{}/mcp", listener.local_addr().unwrap());
2287 let mut server = test_server_config();
2288 server.command = None;
2289 server.url = Some(endpoint.clone());
2290 server.reviewed_plugin = Some(
2291 ReviewedPluginMcpSource::from_authority(
2292 authority,
2293 Some(&endpoint),
2294 Arc::new(crate::plugins::HostEnvironment::default()),
2295 )
2296 .unwrap(),
2297 );
2298
2299 assert_eq!(
2300 oauth::auth_status_for_server("plugin-oauth", &server, None).await,
2301 oauth::McpAuthStatus::Unsupported
2302 );
2303 assert!(
2304 oauth::oauth_login_support(&server, None)
2305 .await
2306 .unwrap()
2307 .is_none()
2308 );
2309 assert!(
2310 oauth::McpOAuthRuntime::from_server_config(
2311 "plugin-oauth",
2312 &server,
2313 reqwest::header::HeaderMap::new(),
2314 )
2315 .await
2316 .unwrap()
2317 .is_none()
2318 );
2319 let login_error =
2320 oauth::perform_oauth_login_for_server("plugin-oauth", &server, None, None, None, None)
2321 .await
2322 .expect_err("plugin OAuth login must be disabled")
2323 .to_string();
2324 assert!(login_error.contains("disabled for plugin-contributed MCP servers"));
2325 let logout_error = oauth::delete_oauth_tokens_for_server("plugin-oauth", &server)
2326 .expect_err("plugin OAuth logout must not touch token storage")
2327 .to_string();
2328 assert!(logout_error.contains("storage is disabled"));
2329
2330 assert!(
2331 tokio::time::timeout(Duration::from_millis(50), listener.accept())
2332 .await
2333 .is_err(),
2334 "plugin OAuth disabled paths must not probe the network"
2335 );
2336 }
2337
2338 fn active_plugin_fixture(
2339 plugin_base: &Path,
2340 ) -> (
2341 crate::plugins::types::LoadedPlugin,
2342 crate::plugins::types::PluginAuthority,
2343 ) {
2344 let plugins_root = plugin_base.parent().expect("plugin parent").to_path_buf();
2345 // Callers use both `<temp>/plugins/<name>` and `<temp>/<name>` layouts.
2346 // Only peel the conventional `plugins` directory; otherwise using the
2347 // parent would place multiple parallel fixtures in the shared system temp
2348 // root and make their durable state files collide on Windows.
2349 let root = if plugins_root.file_name().and_then(|name| name.to_str()) == Some("plugins") {
2350 plugins_root.parent().unwrap_or(&plugins_root).to_path_buf()
2351 } else {
2352 plugins_root.clone()
2353 };
2354 let discovery = crate::plugins::discovery::DiscoveryConfig {
2355 workspace: root.join("project"),
2356 user_plugins_dir: plugins_root,
2357 workspace_plugins_dir: root.join("workspace-plugins-unused"),
2358 builtin_plugin_dirs: Vec::new(),
2359 state_path: root.join("plugin-state").join(format!(
2360 "plugin-state-{}.json",
2361 plugin_base.file_name().unwrap().to_string_lossy()
2362 )),
2363 };
2364 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
2365 let name = registry
2366 .list()
2367 .first()
2368 .expect("discovered plugin")
2369 .name()
2370 .to_string();
2371 registry.trust(&name).unwrap();
2372 registry.enable(&name).unwrap();
2373 (
2374 registry.get(&name).unwrap().clone(),
2375 registry.authority_for(&name).unwrap(),
2376 )
2377 }
2378
2379 #[test]
2380 fn workspace_mcp_config_ignores_project_file_until_workspace_trusted() {
2381 let dir = tempfile::tempdir().unwrap();
2382 let global_path = dir.path().join("global-mcp.json");
2383 let workspace = dir.path().join("workspace");
2384 let project_dir = workspace.join(".codewhale");
2385 let plugin_base = dir.path().join("plugins").join("fixture");
2386 fs::create_dir_all(&project_dir).unwrap();
2387 fs::create_dir_all(&plugin_base).unwrap();
2388 fs::write(
2389 &global_path,
2390 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2391 )
2392 .unwrap();
2393 fs::write(
2394 project_dir.join("mcp.json"),
2395 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2396 )
2397 .unwrap();
2398
2399 let plugins = registry_with_local_mcp("fixture", plugin_base, &workspace);
2400 let cfg = load_config_with_workspace_and_plugins(&global_path, &workspace, &plugins).unwrap();
2401
2402 assert!(cfg.servers.contains_key("global"));
2403 assert!(!cfg.servers.contains_key("project"));
2404 assert!(
2405 cfg.servers
2406 .contains_key(&qualified_plugin_server_name("fixture", "local")),
2407 "user plugin MCP should not be gated by project workspace trust"
2408 );
2409 }
2410
2411 #[test]
2412 fn workspace_mcp_config_ignores_project_local_legacy_trust_marker() {
2413 let dir = tempfile::tempdir().unwrap();
2414 let global_path = dir.path().join("global-mcp.json");
2415 let workspace = dir.path().join("workspace");
2416 let project_dir = workspace.join(".codewhale");
2417 fs::create_dir_all(&project_dir).unwrap();
2418 fs::create_dir_all(workspace.join(".deepseek")).unwrap();
2419 fs::write(workspace.join(".deepseek").join("trusted"), "").unwrap();
2420 fs::write(
2421 &global_path,
2422 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2423 )
2424 .unwrap();
2425 fs::write(
2426 project_dir.join("mcp.json"),
2427 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2428 )
2429 .unwrap();
2430
2431 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
2432
2433 assert!(cfg.servers.contains_key("global"));
2434 assert!(!cfg.servers.contains_key("project"));
2435 }
2436
2437 #[test]
2438 fn workspace_mcp_config_ignores_invalid_untrusted_project_file() {
2439 let dir = tempfile::tempdir().unwrap();
2440 let global_path = dir.path().join("global-mcp.json");
2441 let workspace = dir.path().join("workspace");
2442 let project_dir = workspace.join(".codewhale");
2443 fs::create_dir_all(&project_dir).unwrap();
2444 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
2445 fs::write(project_dir.join("mcp.json"), "{ not json").unwrap();
2446
2447 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
2448
2449 assert!(cfg.servers.is_empty());
2450 }
2451
2452 #[test]
2453 fn workspace_mcp_config_rejects_parent_components() {
2454 let dir = tempfile::tempdir().unwrap();
2455 let global_path = dir.path().join("global-mcp.json");
2456 let workspace = dir.path().join("workspace");
2457 let project_dir = workspace.join(".codewhale");
2458 fs::create_dir_all(&project_dir).unwrap();
2459 let _trust = mark_workspace_trusted(&workspace);
2460 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
2461 fs::write(
2462 project_dir.join("mcp.json"),
2463 r#"{"servers": {"project": {"command": "node", "args": ["server.js"]}}}"#,
2464 )
2465 .unwrap();
2466
2467 let workspace_with_parent = workspace.join("..").join("workspace");
2468 let err = load_config_with_workspace(&global_path, &workspace_with_parent)
2469 .expect_err("parent components in workspace should fail closed");
2470
2471 assert!(
2472 format!("{err:#}").contains("workspace path cannot contain '..'"),
2473 "unexpected error: {err:#}"
2474 );
2475 }
2476
2477 #[test]
2478 fn workspace_mcp_config_resolves_relative_cwd_from_workspace() {
2479 let dir = tempfile::tempdir().unwrap();
2480 let global_path = dir.path().join("global-mcp.json");
2481 let workspace = dir.path().join("workspace");
2482 let project_dir = workspace.join(".codewhale");
2483 fs::create_dir_all(&project_dir).unwrap();
2484 let _trust = mark_workspace_trusted(&workspace);
2485 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
2486 fs::write(
2487 project_dir.join("mcp.json"),
2488 r#"{"servers": {"project": {"command": "node", "args": ["server.js"], "cwd": "tools/mcp"}}}"#,
2489 )
2490 .unwrap();
2491
2492 let cfg = load_config_with_workspace(&global_path, &workspace).unwrap();
2493 let workspace = workspace.canonicalize().unwrap();
2494
2495 let project = cfg.servers.get("project").unwrap();
2496 assert_eq!(
2497 project.cwd.as_deref(),
2498 Some(workspace.join("tools/mcp").as_path())
2499 );
2500 }
2501
2502 #[test]
2503 fn workspace_mcp_config_rejects_project_cwd_escape() {
2504 let dir = tempfile::tempdir().unwrap();
2505 let global_path = dir.path().join("global-mcp.json");
2506 let workspace = dir.path().join("workspace");
2507 let project_dir = workspace.join(".codewhale");
2508 fs::create_dir_all(&project_dir).unwrap();
2509 let _trust = mark_workspace_trusted(&workspace);
2510 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
2511 fs::write(
2512 project_dir.join("mcp.json"),
2513 r#"{"servers": {"project": {"command": "node", "args": ["server.js"], "cwd": "../outside"}}}"#,
2514 )
2515 .unwrap();
2516
2517 let err = load_config_with_workspace(&global_path, &workspace)
2518 .expect_err("project MCP cwd escape must be rejected");
2519
2520 assert!(
2521 err.to_string()
2522 .contains("Project MCP server cwd must stay within workspace"),
2523 "unexpected error: {err}"
2524 );
2525 }
2526
2527 #[cfg(unix)]
2528 #[test]
2529 fn workspace_mcp_config_rejects_symlinked_project_cwd_escape() {
2530 let dir = tempfile::tempdir().unwrap();
2531 let global_path = dir.path().join("global-mcp.json");
2532 let workspace = dir.path().join("workspace");
2533 let project_dir = workspace.join(".codewhale");
2534 let outside = dir.path().join("outside");
2535 fs::create_dir_all(&project_dir).unwrap();
2536 fs::create_dir_all(&outside).unwrap();
2537 std::os::unix::fs::symlink(&outside, workspace.join("tools")).unwrap();
2538 let _trust = mark_workspace_trusted(&workspace);
2539 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
2540 fs::write(
2541 project_dir.join("mcp.json"),
2542 r#"{"servers": {"project": {"command": "node", "args": ["server.js"], "cwd": "tools"}}}"#,
2543 )
2544 .unwrap();
2545
2546 let err = load_config_with_workspace(&global_path, &workspace)
2547 .expect_err("project MCP symlink cwd escape must be rejected");
2548
2549 assert!(
2550 err.to_string()
2551 .contains("Project MCP server cwd must stay within workspace"),
2552 "unexpected error: {err}"
2553 );
2554 }
2555
2556 #[test]
2557 fn workspace_mcp_config_rejects_workspace_traversal() {
2558 let dir = tempfile::tempdir().unwrap();
2559 let global_path = dir.path().join("global-mcp.json");
2560 let workspace = dir.path().join("workspace");
2561 let bad_workspace = workspace.join("..").join("outside");
2562 fs::create_dir_all(&workspace).unwrap();
2563 fs::write(&global_path, r#"{"servers": {}}"#).unwrap();
2564
2565 let err = load_config_with_workspace(&global_path, &bad_workspace)
2566 .expect_err("workspace traversal should fail");
2567 assert!(
2568 format!("{err:#}").contains("workspace path cannot contain '..'"),
2569 "unexpected error: {err:#}"
2570 );
2571 }
2572
2573 #[tokio::test]
2574 async fn workspace_mcp_pool_reload_picks_up_project_config_creation() {
2575 let dir = tempfile::tempdir().unwrap();
2576 let global_path = dir.path().join("global-mcp.json");
2577 let workspace = dir.path().join("workspace");
2578 let project_dir = workspace.join(".codewhale");
2579 fs::create_dir_all(&workspace).unwrap();
2580 let _trust = mark_workspace_trusted(&workspace);
2581 fs::write(
2582 &global_path,
2583 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2584 )
2585 .unwrap();
2586
2587 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
2588 assert_eq!(pool.server_names(), vec!["global".to_string()]);
2589
2590 fs::create_dir_all(&project_dir).unwrap();
2591 fs::write(
2592 project_dir.join("mcp.json"),
2593 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2594 )
2595 .unwrap();
2596
2597 assert!(pool.reload_if_config_changed().await.unwrap());
2598 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
2599 let expected: std::collections::BTreeSet<String> =
2600 ["global".to_string(), "project".to_string()]
2601 .into_iter()
2602 .collect();
2603 assert_eq!(names, expected);
2604 }
2605
2606 #[tokio::test]
2607 async fn workspace_mcp_pool_reload_picks_up_project_config_after_workspace_trust() {
2608 let dir = tempfile::tempdir().unwrap();
2609 let global_path = dir.path().join("global-mcp.json");
2610 let workspace = dir.path().join("workspace");
2611 let project_dir = workspace.join(".codewhale");
2612 fs::create_dir_all(&project_dir).unwrap();
2613 let trust_env = workspace_trust_config_guard(&workspace);
2614 fs::write(
2615 &global_path,
2616 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2617 )
2618 .unwrap();
2619 fs::write(
2620 project_dir.join("mcp.json"),
2621 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2622 )
2623 .unwrap();
2624
2625 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
2626 assert_eq!(pool.server_names(), vec!["global".to_string()]);
2627
2628 write_workspace_trust_config(&trust_env.config_path, &workspace);
2629
2630 assert!(pool.reload_if_config_changed().await.unwrap());
2631 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
2632 let expected: std::collections::BTreeSet<String> =
2633 ["global".to_string(), "project".to_string()]
2634 .into_iter()
2635 .collect();
2636 assert_eq!(names, expected);
2637 }
2638
2639 #[tokio::test]
2640 async fn workspace_mcp_pool_reload_drops_project_config_after_workspace_trust_removed() {
2641 let dir = tempfile::tempdir().unwrap();
2642 let global_path = dir.path().join("global-mcp.json");
2643 let workspace = dir.path().join("workspace");
2644 let project_dir = workspace.join(".codewhale");
2645 fs::create_dir_all(&project_dir).unwrap();
2646 let trust = mark_workspace_trusted(&workspace);
2647 fs::write(
2648 &global_path,
2649 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2650 )
2651 .unwrap();
2652 fs::write(
2653 project_dir.join("mcp.json"),
2654 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2655 )
2656 .unwrap();
2657
2658 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
2659 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
2660 let expected: std::collections::BTreeSet<String> =
2661 ["global".to_string(), "project".to_string()]
2662 .into_iter()
2663 .collect();
2664 assert_eq!(names, expected);
2665
2666 fs::remove_file(&trust.config_path).unwrap();
2667
2668 assert!(pool.reload_if_config_changed().await.unwrap());
2669 assert_eq!(pool.server_names(), vec!["global".to_string()]);
2670 }
2671
2672 #[tokio::test]
2673 async fn workspace_mcp_pool_reload_drops_project_config_after_deletion() {
2674 let dir = tempfile::tempdir().unwrap();
2675 let global_path = dir.path().join("global-mcp.json");
2676 let workspace = dir.path().join("workspace");
2677 let project_dir = workspace.join(".codewhale");
2678 fs::create_dir_all(&project_dir).unwrap();
2679 let _trust = mark_workspace_trusted(&workspace);
2680 fs::write(
2681 &global_path,
2682 r#"{"servers": {"global": {"command": "node", "args": ["global.js"]}}}"#,
2683 )
2684 .unwrap();
2685 let project_path = project_dir.join("mcp.json");
2686 fs::write(
2687 &project_path,
2688 r#"{"servers": {"project": {"command": "php", "args": ["artisan", "boost:mcp"]}}}"#,
2689 )
2690 .unwrap();
2691
2692 let mut pool = McpPool::from_config_path_with_workspace(&global_path, &workspace).unwrap();
2693 let names: std::collections::BTreeSet<String> = pool.server_names().into_iter().collect();
2694 let expected: std::collections::BTreeSet<String> =
2695 ["global".to_string(), "project".to_string()]
2696 .into_iter()
2697 .collect();
2698 assert_eq!(names, expected);
2699
2700 fs::remove_file(project_path).unwrap();
2701
2702 assert!(pool.reload_if_config_changed().await.unwrap());
2703 assert_eq!(pool.server_names(), vec!["global".to_string()]);
2704 }
2705
2706 #[test]
2707 fn test_mcp_config_rejects_traversal_path() {
2708 let err = load_config(Path::new("../mcp.json")).expect_err("traversal path should fail");
2709 assert!(
2710 format!("{err:#}").contains("cannot contain '..'"),
2711 "got: {err:#}"
2712 );
2713 }
2714
2715 #[cfg(unix)]
2716 #[test]
2717 fn mcp_config_rejects_symlinked_config_file() {
2718 let dir = tempfile::tempdir().unwrap();
2719 let target = dir.path().join("target-mcp.json");
2720 let link = dir.path().join("mcp.json");
2721 fs::write(&target, r#"{"servers": {}}"#).expect("write target config");
2722 std::os::unix::fs::symlink(&target, &link).expect("symlink mcp config");
2723
2724 let err = load_config(&link).expect_err("symlinked MCP config should fail");
2725
2726 assert!(format!("{err:#}").contains("regular file"), "got: {err:#}");
2727 }
2728
2729 #[test]
2730 fn init_mcp_config_rejects_traversal_before_parent_creation() {
2731 let dir = tempfile::tempdir().unwrap();
2732 let outside_dir = dir.path().join("outside");
2733 let path = dir
2734 .path()
2735 .join("allowed")
2736 .join("..")
2737 .join("outside")
2738 .join("mcp.json");
2739
2740 let err = init_config(&path, false).expect_err("traversal path should fail");
2741
2742 assert!(
2743 format!("{err:#}").contains("cannot contain '..'"),
2744 "got: {err:#}"
2745 );
2746 assert!(
2747 !outside_dir.exists(),
2748 "init_config must validate before creating parent directories"
2749 );
2750 }
2751
2752 /// A workspace's `.codewhale/mcp.json` is never written through a link: not
2753 /// when `.codewhale` itself is a link, and not when the file is. Reads stay
2754 /// unchanged and the outside directory is never touched.
2755 #[cfg(unix)]
2756 #[test]
2757 fn project_mcp_config_writes_refuse_links_out_of_the_workspace() {
2758 use std::os::unix::fs::symlink;
2759 let workspace = tempfile::tempdir().unwrap();
2760 let outside = tempfile::tempdir().unwrap();
2761 symlink(outside.path(), workspace.path().join(".codewhale")).unwrap();
2762 let path = workspace_mcp_config_path(workspace.path());
2763
2764 let err = init_config(&path, false).expect_err("a linked .codewhale is refused");
2765 assert!(
2766 format!("{err:#}").contains("Refusing symlinked"),
2767 "got: {err:#}"
2768 );
2769 let err = mutate_config(&path, None, |_| Ok(())).expect_err("mutation is refused too");
2770 assert!(
2771 format!("{err:#}").contains("Refusing symlinked"),
2772 "got: {err:#}"
2773 );
2774 assert_eq!(std::fs::read_dir(outside.path()).unwrap().count(), 0);
2775
2776 // A linked file inside a real `.codewhale` is refused as well.
2777 let other = tempfile::tempdir().unwrap();
2778 std::fs::create_dir_all(other.path().join(".codewhale")).unwrap();
2779 let target = outside.path().join("target.json");
2780 symlink(&target, other.path().join(".codewhale").join("mcp.json")).unwrap();
2781 assert!(init_config(&workspace_mcp_config_path(other.path()), false).is_err());
2782 assert!(
2783 !target.exists(),
2784 "a link at the file name must not be created through"
2785 );
2786
2787 // An ordinary workspace still works.
2788 let plain = tempfile::tempdir().unwrap();
2789 assert_eq!(
2790 init_config(&workspace_mcp_config_path(plain.path()), false).unwrap(),
2791 McpWriteStatus::Created
2792 );
2793 }
2794
2795 #[test]
2796 fn test_mcp_config_manager_actions_round_trip() {
2797 let dir = tempfile::tempdir().unwrap();
2798 let path = dir.path().join("mcp.json");
2799
2800 assert_eq!(init_config(&path, false).unwrap(), McpWriteStatus::Created);
2801 assert_eq!(
2802 init_config(&path, false).unwrap(),
2803 McpWriteStatus::SkippedExists
2804 );
2805
2806 add_server_config(
2807 &path,
2808 "local".to_string(),
2809 Some("node".to_string()),
2810 None,
2811 vec!["server.js".to_string()],
2812 None,
2813 )
2814 .unwrap();
2815 set_server_enabled(&path, "local", false).unwrap();
2816 let disabled = manager_snapshot_from_config(&path, true).unwrap();
2817 let local = disabled
2818 .servers
2819 .iter()
2820 .find(|server| server.name == "local")
2821 .unwrap();
2822 assert!(!local.enabled);
2823 assert_eq!(local.transport, "stdio");
2824
2825 remove_server_config(&path, "local").unwrap();
2826 let removed = manager_snapshot_from_config(&path, true).unwrap();
2827 assert!(removed.servers.iter().all(|server| server.name != "local"));
2828 }
2829
2830 #[test]
2831 fn test_mcp_config_adds_explicit_sse_transport() {
2832 let dir = tempfile::tempdir().unwrap();
2833 let path = dir.path().join("mcp.json");
2834
2835 add_server_config(
2836 &path,
2837 "legacy".to_string(),
2838 None,
2839 Some("https://example.com/v1/mcp/sse".to_string()),
2840 Vec::new(),
2841 Some("sse".to_string()),
2842 )
2843 .unwrap();
2844
2845 let cfg = load_config(&path).unwrap();
2846 assert_eq!(
2847 cfg.servers
2848 .get("legacy")
2849 .and_then(|server| server.transport.as_deref()),
2850 Some("sse")
2851 );
2852
2853 let snapshot = manager_snapshot_from_config(&path, false).unwrap();
2854 assert_eq!(snapshot.servers[0].transport, "sse");
2855 }
2856
2857 #[test]
2858 fn test_mcp_config_rejects_unknown_transport() {
2859 let dir = tempfile::tempdir().unwrap();
2860 let path = dir.path().join("mcp.json");
2861
2862 let err = add_server_config(
2863 &path,
2864 "bad".to_string(),
2865 None,
2866 Some("https://example.com/mcp".to_string()),
2867 Vec::new(),
2868 Some("streamable".to_string()),
2869 )
2870 .expect_err("unknown transport should fail");
2871
2872 assert!(
2873 format!("{err:#}").contains("Unsupported MCP transport"),
2874 "got: {err:#}"
2875 );
2876 }
2877
2878 #[test]
2879 fn test_server_effective_timeouts() {
2880 let global = McpTimeouts::default();
2881
2882 let server_with_override = McpServerConfig {
2883 command: Some("test".to_string()),
2884 args: vec![],
2885 env: HashMap::new(),
2886 cwd: None,
2887 url: None,
2888 transport: None,
2889 connect_timeout: Some(20),
2890 execute_timeout: None,
2891 read_timeout: Some(180),
2892 disabled: false,
2893 enabled: true,
2894 required: false,
2895 enabled_tools: Vec::new(),
2896 disabled_tools: Vec::new(),
2897 headers: HashMap::new(),
2898 env_headers: HashMap::new(),
2899 bearer_token_env_var: None,
2900 scopes: Vec::new(),
2901 oauth: None,
2902 oauth_resource: None,
2903 reviewed_plugin: None,
2904 runtime_added: false,
2905 allow_private_network: false,
2906 };
2907
2908 assert_eq!(server_with_override.effective_connect_timeout(&global), 20);
2909 assert_eq!(
2910 server_with_override.effective_execute_timeout(&global),
2911 1800
2912 ); // global default
2913 assert_eq!(server_with_override.effective_read_timeout(&global), 180);
2914 }
2915
2916 #[test]
2917 fn test_mcp_pool_is_mcp_tool() {
2918 assert!(McpPool::is_mcp_tool("mcp_filesystem_read"));
2919 assert!(McpPool::is_mcp_tool("mcp_git_status"));
2920 assert!(McpPool::is_mcp_tool("list_mcp_resources"));
2921 assert!(McpPool::is_mcp_tool("list_mcp_resource_templates"));
2922 assert!(McpPool::is_mcp_tool("read_mcp_resource"));
2923 assert!(!McpPool::is_mcp_tool("read_file"));
2924 assert!(!McpPool::is_mcp_tool("exec_shell"));
2925 }
2926
2927 struct ScriptedValueTransport {
2928 sent: Arc<Mutex<Vec<serde_json::Value>>>,
2929 responses: VecDeque<Vec<u8>>,
2930 }
2931
2932 #[async_trait::async_trait]
2933 impl McpTransport for ScriptedValueTransport {
2934 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
2935 self.sent
2936 .lock()
2937 .unwrap()
2938 .push(serde_json::from_slice(&msg)?);
2939 Ok(())
2940 }
2941
2942 async fn recv(&mut self) -> Result<Vec<u8>> {
2943 self.responses
2944 .pop_front()
2945 .context("scripted transport exhausted")
2946 }
2947 }
2948
2949 struct HangingValueTransport {
2950 sent: Arc<Mutex<Vec<serde_json::Value>>>,
2951 }
2952
2953 #[async_trait::async_trait]
2954 impl McpTransport for HangingValueTransport {
2955 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
2956 self.sent
2957 .lock()
2958 .unwrap()
2959 .push(serde_json::from_slice(&msg)?);
2960 Ok(())
2961 }
2962
2963 async fn recv(&mut self) -> Result<Vec<u8>> {
2964 std::future::pending().await
2965 }
2966 }
2967
2968 struct ScriptedThenHangingTransport {
2969 sent: Arc<Mutex<Vec<serde_json::Value>>>,
2970 responses: VecDeque<Vec<u8>>,
2971 }
2972
2973 #[async_trait::async_trait]
2974 impl McpTransport for ScriptedThenHangingTransport {
2975 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
2976 self.sent
2977 .lock()
2978 .unwrap()
2979 .push(serde_json::from_slice(&msg)?);
2980 Ok(())
2981 }
2982
2983 async fn recv(&mut self) -> Result<Vec<u8>> {
2984 match self.responses.pop_front() {
2985 Some(response) => Ok(response),
2986 None => std::future::pending().await,
2987 }
2988 }
2989 }
2990
2991 /// A transport that answers inside `send`, as Streamable HTTP reads the reply
2992 /// within the POST, and never finishes that send.
2993 struct HangingSendTransport;
2994
2995 #[async_trait::async_trait]
2996 impl McpTransport for HangingSendTransport {
2997 async fn send(&mut self, _msg: Vec<u8>) -> Result<()> {
2998 std::future::pending().await
2999 }
3000
3001 async fn recv(&mut self) -> Result<Vec<u8>> {
3002 std::future::pending().await
3003 }
3004 }
3005
3006 /// A transport whose write side is gone — the shape a crashed or exited
3007 /// stdio MCP child leaves behind (EPIPE on the next `write_all`).
3008 struct FailingSendTransport;
3009
3010 #[async_trait::async_trait]
3011 impl McpTransport for FailingSendTransport {
3012 async fn send(&mut self, _msg: Vec<u8>) -> Result<()> {
3013 anyhow::bail!("Broken pipe (os error 32)")
3014 }
3015
3016 async fn recv(&mut self) -> Result<Vec<u8>> {
3017 std::future::pending().await
3018 }
3019 }
3020
3021 struct DropCountingTransport {
3022 drops: Arc<AtomicUsize>,
3023 }
3024
3025 #[async_trait::async_trait]
3026 impl McpTransport for DropCountingTransport {
3027 async fn send(&mut self, _msg: Vec<u8>) -> Result<()> {
3028 Ok(())
3029 }
3030
3031 async fn recv(&mut self) -> Result<Vec<u8>> {
3032 std::future::pending().await
3033 }
3034 }
3035
3036 impl Drop for DropCountingTransport {
3037 fn drop(&mut self) {
3038 self.drops.fetch_add(1, AtomicOrdering::SeqCst);
3039 }
3040 }
3041
3042 fn test_server_config() -> McpServerConfig {
3043 McpServerConfig {
3044 command: Some("mock".to_string()),
3045 args: Vec::new(),
3046 env: HashMap::new(),
3047 cwd: None,
3048 url: None,
3049 transport: None,
3050 connect_timeout: None,
3051 execute_timeout: None,
3052 read_timeout: None,
3053 disabled: false,
3054 enabled: true,
3055 required: false,
3056 enabled_tools: Vec::new(),
3057 disabled_tools: Vec::new(),
3058 headers: HashMap::new(),
3059 env_headers: HashMap::new(),
3060 bearer_token_env_var: None,
3061 scopes: Vec::new(),
3062 oauth: None,
3063 oauth_resource: None,
3064 reviewed_plugin: None,
3065 runtime_added: false,
3066 allow_private_network: false,
3067 }
3068 }
3069
3070 fn test_connection(transport: Box<dyn McpTransport>) -> McpConnection {
3071 McpConnection {
3072 name: "mock".to_string(),
3073 transport,
3074 tools: Vec::new(),
3075 resources: Vec::new(),
3076 resource_templates: Vec::new(),
3077 prompts: Vec::new(),
3078 request_id: AtomicU64::new(1),
3079 state: ConnectionState::Ready,
3080 config: test_server_config(),
3081 server_capabilities: None,
3082 instructions: None,
3083 discovery_timeout: Duration::from_secs(default_connect_timeout()),
3084 read_timeout_secs: default_read_timeout(),
3085 cancel_token: tokio_util::sync::CancellationToken::new(),
3086 authority_revocation_reason: Arc::new(std::sync::Mutex::new(None)),
3087 authority_watch: None,
3088 catalog_generation: 0,
3089 decision_key: None,
3090 }
3091 }
3092
3093 #[cfg(unix)]
3094 #[tokio::test]
3095 async fn execute_timeout_after_partial_stdio_response_does_not_corrupt_next_call() -> Result<()> {
3096 use serde_json::{Value, json};
3097 struct SharedStdio(Arc<tokio::sync::Mutex<StdioTransport>>);
3098 #[async_trait::async_trait]
3099 impl McpTransport for SharedStdio {
3100 async fn send(&mut self, bytes: Vec<u8>) -> Result<()> {
3101 self.0.lock().await.send(bytes).await
3102 }
3103 async fn recv(&mut self) -> Result<Vec<u8>> {
3104 self.0.lock().await.recv().await
3105 }
3106 }
3107 let dir = tempfile::tempdir()?;
3108 let requests = dir.path().join("requests.jsonl");
3109 let ready = dir.path().join("ready");
3110 let script = r#"
3111 printf '%s' '{"jsonrpc":"2.0","id":"1","result":'
3112 : > "$2"
3113 IFS= read -r first
3114 printf '%s\n' "$first" >> "$1"
3115 IFS= read -r second
3116 printf '%s\n' "$second" >> "$1"
3117 printf '%s\n' 'null}' '{"jsonrpc":"2.0","id":"2","result":{"ok":true}}'
3118 IFS= read -r keep_open
3119 "#;
3120 let mut config = test_server_config();
3121 config.args = vec![
3122 "-c".into(),
3123 script.into(),
3124 "cw-partial-frame-fixture".into(),
3125 requests.display().to_string(),
3126 ready.display().to_string(),
3127 ];
3128 let transport = Arc::new(tokio::sync::Mutex::new(StdioTransport::spawn(
3129 "partial-frame",
3130 "sh",
3131 &config,
3132 tokio_util::sync::CancellationToken::new(),
3133 )?));
3134 tokio::time::timeout(Duration::from_secs(10), async {
3135 while !ready.exists() {
3136 tokio::time::sleep(Duration::from_millis(5)).await;
3137 }
3138 })
3139 .await?;
3140 let mut connection = test_connection(Box::new(SharedStdio(Arc::clone(&transport))));
3141 let error = connection
3142 .call_tool("first", json!({}), 1)
3143 .await
3144 .unwrap_err();
3145 assert!(error.to_string().contains("timed out"), "{error:#}");
3146 assert_eq!(
3147 transport.lock().await.pending_line,
3148 br#"{"jsonrpc":"2.0","id":"1","result":"#
3149 );
3150 assert!(connection.is_ready());
3151 assert_eq!(
3152 connection.call_tool("second", json!({}), 5).await?,
3153 json!({"ok": true})
3154 );
3155 let sent = fs::read_to_string(requests)?;
3156 let sent = sent
3157 .lines()
3158 .map(serde_json::from_str::<Value>)
3159 .collect::<std::result::Result<Vec<_>, _>>()?;
3160 assert_eq!(sent.len(), 2, "neither request may be replayed");
3161 assert_eq!(sent[0]["id"], "1");
3162 assert_eq!(sent[1]["id"], "2");
3163 transport.lock().await.shutdown().await;
3164 Ok::<_, anyhow::Error>(())
3165 }
3166
3167 fn json_frame(value: serde_json::Value) -> Vec<u8> {
3168 serde_json::to_vec(&value).unwrap()
3169 }
3170
3171 #[test]
3172 fn http_request_ceiling_covers_the_execute_budget() {
3173 let mut config = test_server_config();
3174 config.execute_timeout = Some(1800);
3175 config.read_timeout = Some(120);
3176 let global = McpTimeouts {
3177 connect_timeout: 10,
3178 execute_timeout: 1800,
3179 read_timeout: 120,
3180 };
3181 assert_eq!(http_request_ceiling_secs(&config, &global), 1800);
3182
3183 // A per-server read knob above both stays intact.
3184 config.read_timeout = Some(3600);
3185 assert_eq!(http_request_ceiling_secs(&config, &global), 3600);
3186 }
3187
3188 /// A transport that stays silent for a fixed delay, then answers with a
3189 /// matching result — the shape of an MCP server executing a long tool call.
3190 struct DelayedResponseTransport {
3191 delay: Duration,
3192 pending: Option<Vec<u8>>,
3193 }
3194
3195 #[async_trait::async_trait]
3196 impl McpTransport for DelayedResponseTransport {
3197 async fn send(&mut self, msg: Vec<u8>) -> Result<()> {
3198 // Echo the request id so every call gets its matching response.
3199 let request: serde_json::Value = serde_json::from_slice(&msg)?;
3200 let response = serde_json::json!({
3201 "jsonrpc": "2.0",
3202 "id": request["id"].clone(),
3203 "result": {"ok": true}
3204 });
3205 self.pending = Some(json_frame(response));
3206 Ok(())
3207 }
3208
3209 async fn recv(&mut self) -> Result<Vec<u8>> {
3210 tokio::time::sleep(self.delay).await;
3211 self.pending.take().context("delayed transport exhausted")
3212 }
3213 }
3214
3215 #[tokio::test]
3216 async fn a_tool_response_after_the_read_knob_still_completes_within_the_execute_budget() {
3217 let mut connection = test_connection(Box::new(DelayedResponseTransport {
3218 // The response arrives after the read knob (1s) but well inside the
3219 // execute budget (30s): the read knob must not bound a request's
3220 // receive and mark the connection dead before the reply lands.
3221 delay: Duration::from_millis(1500),
3222 pending: None,
3223 }));
3224 connection.read_timeout_secs = 1;
3225 let result = connection
3226 .call_tool("slow", serde_json::json!({}), 30)
3227 .await
3228 .expect("a silent tool execution must not be cut off by the read knob");
3229 assert_eq!(result, serde_json::json!({"ok": true}));
3230 assert!(
3231 connection.is_ready(),
3232 "a completed call keeps the connection"
3233 );
3234 }
3235
3236 /// A wedged server fails a request at the request's own budget, and the
3237 /// request is abandoned, not the connection: its late reply carries the
3238 /// abandoned id and is skipped (see the stdio fixtures below). The read knob
3239 /// (1s) and the budget (2s) are distinct on purpose: equal deadlines used to
3240 /// race, and whichever timer won decided whether the connection survived.
3241 /// The per-frame read-knob disconnect is still pinned for the handshake by
3242 /// `recv_times_out_waiting_for_mcp_response_and_disconnects`.
3243 #[tokio::test]
3244 async fn a_wedged_request_fails_at_its_own_budget_and_keeps_the_connection() {
3245 let mut connection = test_connection(Box::new(HangingValueTransport {
3246 sent: Arc::new(Mutex::new(Vec::new())),
3247 }));
3248 connection.read_timeout_secs = 1;
3249 let error = connection
3250 .read_resource("file:///wedged", 2)
3251 .await
3252 .expect_err("a wedged server must fail the request at its budget");
3253 assert!(
3254 error
3255 .to_string()
3256 .contains("MCP method 'resources/read' on server 'mock' timed out after 2s"),
3257 "the request budget, not the read knob, must end the request: {error:#}"
3258 );
3259 assert!(
3260 connection.is_ready(),
3261 "an expired request must not declare the connection dead"
3262 );
3263 }
3264
3265 /// A request whose transport blocks inside `send` (Streamable HTTP) still
3266 /// ends at its own budget, not at the transport's larger client ceiling, and
3267 /// the connection is rebuilt because the abandoned write may be partial.
3268 #[tokio::test]
3269 async fn a_request_blocked_inside_send_ends_at_its_own_budget() {
3270 let mut connection = test_connection(Box::new(HangingSendTransport));
3271 connection.read_timeout_secs = 30;
3272 let error = tokio::time::timeout(
3273 std::time::Duration::from_secs(10),
3274 connection.read_resource("file:///wedged-post", 1),
3275 )
3276 .await
3277 .expect("the request budget, not the transport, must end a blocked send")
3278 .expect_err("a POST that never completes must end at the request budget");
3279 assert!(
3280 error
3281 .to_string()
3282 .contains("MCP method 'resources/read' on server 'mock' timed out after 1s"),
3283 "{error:#}"
3284 );
3285 assert!(
3286 !connection.is_ready(),
3287 "a send abandoned mid-write leaves the frame boundary unknown, so the connection is rebuilt"
3288 );
3289 }
3290
3291 #[tokio::test]
3292 async fn call_method_skips_notifications_and_unmatched_responses() {
3293 let sent = Arc::new(Mutex::new(Vec::new()));
3294 let transport = ScriptedValueTransport {
3295 sent: Arc::clone(&sent),
3296 responses: VecDeque::from([
3297 json_frame(serde_json::json!({
3298 "jsonrpc": "2.0",
3299 "method": "notifications/progress",
3300 "params": {"progress": 0.5}
3301 })),
3302 json_frame(serde_json::json!({
3303 "jsonrpc": "2.0",
3304 "id": 99,
3305 "result": {"ignored": true}
3306 })),
3307 json_frame(serde_json::json!({
3308 "jsonrpc": "2.0",
3309 "id": 1,
3310 "result": {"ok": true}
3311 })),
3312 ]),
3313 };
3314 let mut conn = test_connection(Box::new(transport));
3315
3316 let result = conn
3317 .call_method("tools/call", serde_json::json!({"name": "echo"}), 1)
3318 .await
3319 .unwrap();
3320
3321 assert_eq!(result, serde_json::json!({"ok": true}));
3322 let sent = sent.lock().unwrap();
3323 assert_eq!(sent.len(), 1);
3324 assert_eq!(sent[0]["jsonrpc"], "2.0");
3325 assert_eq!(sent[0]["id"], "1");
3326 assert_eq!(sent[0]["method"], "tools/call");
3327 }
3328
3329 #[tokio::test]
3330 async fn call_method_invalid_json_includes_server_output_preview() {
3331 let sent = Arc::new(Mutex::new(Vec::new()));
3332 let transport = ScriptedValueTransport {
3333 sent: Arc::clone(&sent),
3334 responses: VecDeque::from([b"Allow Burp MCP connection? [y/N]".to_vec()]),
3335 };
3336 let mut conn = test_connection(Box::new(transport));
3337
3338 let err = conn
3339 .call_method("tools/call", serde_json::json!({"name": "burp"}), 1)
3340 .await
3341 .expect_err("non-json MCP stdout should fail");
3342 let msg = err.to_string();
3343
3344 assert!(msg.contains("Invalid MCP JSON-RPC message from server 'mock'"));
3345 assert!(msg.contains("Allow Burp MCP connection"));
3346 assert_eq!(conn.state(), ConnectionState::Disconnected);
3347 }
3348
3349 #[tokio::test]
3350 async fn recv_times_out_waiting_for_mcp_response_and_disconnects() {
3351 let sent = Arc::new(Mutex::new(Vec::new()));
3352 let mut conn = test_connection(Box::new(HangingValueTransport {
3353 sent: Arc::clone(&sent),
3354 }));
3355 conn.read_timeout_secs = 0;
3356
3357 let err = conn
3358 .recv("1".to_string())
3359 .await
3360 .expect_err("hung transport should time out inside recv");
3361
3362 assert!(
3363 err.to_string()
3364 .contains("Timed out waiting for MCP JSON-RPC response from server 'mock' after 0s"),
3365 "unexpected error: {err:#}"
3366 );
3367 assert_eq!(conn.state(), ConnectionState::Disconnected);
3368 }
3369
3370 #[tokio::test]
3371 async fn call_method_times_out_while_waiting_for_response() {
3372 let sent = Arc::new(Mutex::new(Vec::new()));
3373 let mut conn = test_connection(Box::new(HangingValueTransport {
3374 sent: Arc::clone(&sent),
3375 }));
3376
3377 let err = conn
3378 .call_method("tools/call", serde_json::json!({"name": "echo"}), 0)
3379 .await
3380 .expect_err("hung receive should time out");
3381
3382 assert!(
3383 err.to_string()
3384 .contains("MCP method 'tools/call' on server 'mock' timed out after 0s"),
3385 "unexpected error: {err:#}"
3386 );
3387 assert_eq!(sent.lock().unwrap().len(), 1);
3388 }
3389
3390 /// JSON-RPC requires exactly one of `result` / `error` on a response. A
3391 /// response carrying neither is a broken server, and reporting it as a
3392 /// successful call with a `null` payload is a fake success: the tool result
3393 /// reaches the model as `ToolResult::success("null")`, indistinguishable from
3394 /// a tool that genuinely did nothing.
3395 #[tokio::test]
3396 async fn call_method_rejects_a_response_with_neither_result_nor_error() {
3397 let sent = Arc::new(Mutex::new(Vec::new()));
3398 let transport = ScriptedValueTransport {
3399 sent: Arc::clone(&sent),
3400 responses: VecDeque::from([json_frame(serde_json::json!({
3401 "jsonrpc": "2.0",
3402 "id": 1
3403 }))]),
3404 };
3405 let mut conn = test_connection(Box::new(transport));
3406
3407 let err = conn
3408 .call_method("tools/call", serde_json::json!({"name": "echo"}), 1)
3409 .await
3410 .expect_err("a result-less, error-less response is not a successful call");
3411 let rendered = format!("{err:#}");
3412 assert!(
3413 rendered.contains("neither a result nor an error"),
3414 "unexpected error: {rendered}"
3415 );
3416 assert!(
3417 rendered.contains("tools/call"),
3418 "unexpected error: {rendered}"
3419 );
3420 }
3421
3422 /// …while an *explicit* `"result": null` is a well-formed empty success and
3423 /// must keep flowing through unchanged.
3424 #[tokio::test]
3425 async fn call_method_preserves_an_explicit_null_result() {
3426 let sent = Arc::new(Mutex::new(Vec::new()));
3427 let transport = ScriptedValueTransport {
3428 sent: Arc::clone(&sent),
3429 responses: VecDeque::from([json_frame(serde_json::json!({
3430 "jsonrpc": "2.0",
3431 "id": 1,
3432 "result": null
3433 }))]),
3434 };
3435 let mut conn = test_connection(Box::new(transport));
3436
3437 let result = conn
3438 .call_method("tools/call", serde_json::json!({"name": "echo"}), 1)
3439 .await
3440 .expect("an explicit null result is a valid response");
3441 assert_eq!(result, serde_json::Value::Null);
3442 }
3443
3444 /// A failed *write* has to disconnect the connection, exactly like a failed
3445 /// read does. `McpPool::get_or_connect` reuses any connection whose
3446 /// `is_ready()` is true, so a connection left in `Ready` after its transport
3447 /// write side died is never rebuilt — the pool hands the same dead child back
3448 /// on every later tool call and the server stays broken for the rest of the
3449 /// session even though a reconnect would fix it.
3450 #[tokio::test]
3451 async fn call_method_disconnects_when_the_transport_write_side_is_gone() {
3452 let mut conn = test_connection(Box::new(FailingSendTransport));
3453
3454 let err = conn
3455 .call_method("tools/call", serde_json::json!({"name": "echo"}), 1)
3456 .await
3457 .expect_err("a dead write side must fail the call");
3458 assert!(
3459 format!("{err:#}").contains("Broken pipe"),
3460 "unexpected error: {err:#}"
3461 );
3462 assert_eq!(conn.state(), ConnectionState::Disconnected);
3463 assert!(
3464 !conn.is_ready(),
3465 "a connection whose write side died must not be reused"
3466 );
3467 }
3468
3469 /// The pool-level consequence of the same defect: `/mcp` (and every
3470 /// `connected_servers` caller) reported a server with a dead write side as
3471 /// still connected, and `get_or_connect` handed the dead connection back
3472 /// instead of rebuilding it.
3473 #[tokio::test]
3474 async fn pool_stops_advertising_a_server_whose_write_side_died() {
3475 let dir = tempfile::tempdir().unwrap();
3476 let path = dir.path().join("mcp.json");
3477 fs::write(
3478 &path,
3479 r#"{
3480 "mcpServers": {
3481 "mock": {
3482 "command": "codewhale-tui-test-this-binary-does-not-exist-9f8e7d6c5b4a",
3483 "args": []
3484 }
3485 }
3486 }"#,
3487 )
3488 .unwrap();
3489 let mut pool = McpPool::from_config_path(&path).unwrap();
3490 let mut conn = test_connection(Box::new(FailingSendTransport));
3491 conn.tools.push(McpTool {
3492 name: "echo".to_string(),
3493 description: None,
3494 input_schema: serde_json::json!({"type": "object"}),
3495 annotations: None,
3496 });
3497 pool.connections.insert("mock".to_string(), conn);
3498 assert_eq!(pool.connected_servers(), vec!["mock"]);
3499
3500 let err = pool
3501 .call_tool("mcp_mock_echo", serde_json::json!({}))
3502 .await
3503 .expect_err("a dead write side must fail the call");
3504 assert!(
3505 format!("{err:#}").contains("Broken pipe"),
3506 "unexpected error: {err:#}"
3507 );
3508
3509 assert!(
3510 pool.connected_servers().is_empty(),
3511 "a server whose write side died must not report as connected"
3512 );
3513 let reconnect = match pool.get_or_connect("mock").await {
3514 Ok(_) => panic!("the pool must rebuild rather than reuse the dead connection"),
3515 Err(error) => error,
3516 };
3517 assert!(
3518 format!("{reconnect:#}").contains("spawn failed"),
3519 "expected a fresh spawn attempt, got: {reconnect:#}"
3520 );
3521 }
3522
3523 /// #6187: a failed reconnect must not erase the previous connection — the
3524 /// last-good tool catalog stays registered (model-visible, since catalog
3525 /// aggregation filters on authority, not liveness) for the whole outage,
3526 /// while the restored connection stays non-ready so `get_or_connect`
3527 /// keeps retrying per the backoff.
3528 #[tokio::test]
3529 async fn failed_reconnect_restores_last_good_catalog() {
3530 let dir = tempfile::tempdir().unwrap();
3531 let path = dir.path().join("mcp.json");
3532 fs::write(
3533 &path,
3534 r#"{
3535 "mcpServers": {
3536 "mock": {
3537 "command": "codewhale-tui-test-this-binary-does-not-exist-9f8e7d6c5b4a",
3538 "args": []
3539 }
3540 }
3541 }"#,
3542 )
3543 .unwrap();
3544 let mut pool = McpPool::from_config_path(&path).unwrap();
3545 let mut conn = test_connection(Box::new(HangingValueTransport {
3546 sent: Arc::new(Mutex::new(Vec::new())),
3547 }));
3548 conn.name = "mock".to_string();
3549 conn.config = pool.config.servers.get("mock").unwrap().clone();
3550 conn.catalog_generation = pool.current_catalog_generation();
3551 // The shape a crashed server leaves behind: not ready, but its
3552 // last-good catalog is still discovered on the connection.
3553 conn.state = ConnectionState::Disconnected;
3554 conn.tools.push(McpTool {
3555 name: "echo".to_string(),
3556 description: None,
3557 input_schema: serde_json::json!({"type": "object"}),
3558 annotations: None,
3559 });
3560 pool.connections.insert("mock".to_string(), conn);
3561
3562 // `&mut McpConnection` is not `Debug`, so mirror the sibling test's
3563 // match instead of `expect_err`.
3564 let error = match pool.get_or_connect("mock").await {
3565 Ok(_) => panic!("reconnect against a missing binary must fail"),
3566 Err(error) => error,
3567 };
3568 assert!(
3569 format!("{error:#}").contains("spawn failed"),
3570 "unexpected error: {error:#}"
3571 );
3572
3573 let restored = pool
3574 .connections
3575 .get("mock")
3576 .expect("failed reconnect must restore the previous connection");
3577 assert!(
3578 !restored.is_ready(),
3579 "the restored connection must stay non-ready so the pool keeps retrying"
3580 );
3581 assert!(
3582 pool.all_tools()
3583 .iter()
3584 .any(|(name, _)| name == "mcp_mock_echo"),
3585 "the model-visible tool surface must survive the failed reconnect"
3586 );
3587 assert_eq!(
3588 restored.tools.len(),
3589 1,
3590 "the restored connection must keep its last-good catalog"
3591 );
3592 }
3593
3594 #[tokio::test]
3595 async fn test_mcp_pool_empty_config() {
3596 let pool = McpPool::new(McpConfig::default());
3597 assert!(pool.server_names().is_empty());
3598 assert!(pool.all_tools().is_empty());
3599 }
3600
3601 /// #1267 part 2: a pool built without a source path has no file to watch,
3602 /// so `reload_if_config_changed` must short-circuit instead of trying
3603 /// to stat `/`.
3604 #[tokio::test]
3605 async fn reload_if_config_changed_is_noop_without_source_path() {
3606 let mut pool = McpPool::new(McpConfig::default());
3607 let reloaded = pool.reload_if_config_changed().await.unwrap();
3608 assert!(!reloaded, "no source path → no reload");
3609 }
3610
3611 /// #1267 part 2: when the on-disk config is byte-unchanged, the lazy
3612 /// reload must not drop connections — every call to `get_or_connect`
3613 /// would otherwise pay a full reconnect cycle on networked filesystems
3614 /// where mtime granularity is coarse.
3615 #[tokio::test]
3616 async fn reload_if_config_changed_skips_when_content_unchanged() {
3617 let dir = tempfile::tempdir().unwrap();
3618 let path = dir.path().join("mcp.json");
3619 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
3620 let mut pool = McpPool::from_config_path(&path).unwrap();
3621 // Force the mtime to advance without changing content.
3622 std::thread::sleep(std::time::Duration::from_millis(10));
3623 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
3624 let reloaded = pool.reload_if_config_changed().await.unwrap();
3625 assert!(
3626 !reloaded,
3627 "content-unchanged config must not trigger a reload"
3628 );
3629 }
3630
3631 /// #1267 part 2: when the on-disk config changes content, the next
3632 /// `reload_if_config_changed` call must swap in the new config and
3633 /// (would) drop all live connections. We can't stand up a real
3634 /// `McpConnection` in a unit test, so we observe the swap via the
3635 /// publicly-readable side: server names go from empty to non-empty.
3636 #[tokio::test]
3637 async fn reload_if_config_changed_swaps_config_on_content_change() {
3638 let dir = tempfile::tempdir().unwrap();
3639 let path = dir.path().join("mcp.json");
3640 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
3641 let mut pool = McpPool::from_config_path(&path).unwrap();
3642 assert!(pool.server_names().is_empty());
3643 // Mutate the file so both the mtime and the hash change.
3644 std::thread::sleep(std::time::Duration::from_millis(10));
3645 std::fs::write(
3646 &path,
3647 r#"{"servers":{"new":{"command":"echo","args":["hi"]}}}"#,
3648 )
3649 .unwrap();
3650 let reloaded = pool.reload_if_config_changed().await.unwrap();
3651 assert!(reloaded, "content-changed config must trigger reload");
3652 let names = pool.server_names();
3653 assert!(
3654 names.contains(&"new".to_string()),
3655 "expected new server in pool after reload, got {names:?}"
3656 );
3657 }
3658
3659 #[tokio::test]
3660 async fn stale_handshake_cannot_be_restamped_after_config_reload() {
3661 let dir = tempfile::tempdir().unwrap();
3662 let path = dir.path().join("mcp.json");
3663 std::fs::write(&path, r#"{"servers":{"local":{"command":"node"}}}"#).unwrap();
3664 let mut pool = McpPool::from_config_path(&path).unwrap();
3665 let drops = Arc::new(AtomicUsize::new(0));
3666 let mut connection = test_connection(Box::new(DropCountingTransport {
3667 drops: drops.clone(),
3668 }));
3669 connection.catalog_generation = pool.current_catalog_generation();
3670 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
3671 pool.reload_from_config_sources(true).unwrap();
3672 let error = pool
3673 .store_ready_connection("local".to_string(), connection)
3674 .unwrap_err();
3675 assert!(error.to_string().contains("configuration changed"));
3676 assert!(!pool.connections.contains_key("local"));
3677 assert_eq!(drops.load(AtomicOrdering::SeqCst), 1);
3678 }
3679
3680 #[tokio::test]
3681 async fn reload_if_config_changed_drops_live_connections() {
3682 let dir = tempfile::tempdir().unwrap();
3683 let path = dir.path().join("mcp.json");
3684 std::fs::write(
3685 &path,
3686 r#"{"servers":{"local":{"command":"node","args":["server.js"]}}}"#,
3687 )
3688 .unwrap();
3689 let mut pool = McpPool::from_config_path(&path).unwrap();
3690 let drops = Arc::new(AtomicUsize::new(0));
3691 let mut conn = test_connection(Box::new(DropCountingTransport {
3692 drops: Arc::clone(&drops),
3693 }));
3694 conn.name = "local".to_string();
3695 conn.config = pool.config.servers.get("local").unwrap().clone();
3696 pool.connections.insert("local".to_string(), conn);
3697
3698 std::thread::sleep(std::time::Duration::from_millis(10));
3699 std::fs::write(
3700 &path,
3701 r#"{"servers":{"local":{"command":"node","args":["server-v2.js"]}}}"#,
3702 )
3703 .unwrap();
3704
3705 let reloaded = pool.reload_if_config_changed().await.unwrap();
3706 assert!(reloaded, "content-changed config must trigger reload");
3707 assert_eq!(
3708 drops.load(AtomicOrdering::SeqCst),
3709 1,
3710 "reload must drop the stale live transport"
3711 );
3712 assert!(
3713 !pool.connections.contains_key("local"),
3714 "stale connection must not survive config reload"
3715 );
3716 assert_eq!(
3717 pool.config.servers.get("local").unwrap().args,
3718 vec!["server-v2.js".to_string()]
3719 );
3720 }
3721
3722 #[tokio::test]
3723 async fn connect_all_reloads_before_snapshotting_new_server_names() {
3724 let dir = tempfile::tempdir().unwrap();
3725 let path = dir.path().join("mcp.json");
3726 std::fs::write(&path, r#"{"servers":{}}"#).unwrap();
3727 let mut pool = McpPool::from_config_path(&path).unwrap();
3728
3729 std::fs::write(
3730 &path,
3731 r#"{"servers":{"late":{"command":"codewhale-test-command-that-does-not-exist"}}}"#,
3732 )
3733 .unwrap();
3734 // Make the test independent of filesystem mtime granularity.
3735 pool.last_mtimes = vec![None];
3736
3737 let errors = pool.connect_all().await;
3738 assert!(
3739 pool.server_names().contains(&"late".to_string()),
3740 "the first connect_all call must install the changed config"
3741 );
3742 assert!(
3743 errors.iter().any(|(name, _)| name == "late"),
3744 "the newly-added server must be attempted on the same call"
3745 );
3746 }
3747
3748 #[tokio::test]
3749 async fn explicit_reload_reconnects_unchanged_config_and_preserves_dynamic_servers() {
3750 let dir = tempfile::tempdir().unwrap();
3751 let path = dir.path().join("mcp.json");
3752 std::fs::write(
3753 &path,
3754 r#"{"servers":{"local":{"command":"node","disabled":true}}}"#,
3755 )
3756 .unwrap();
3757 let mut pool = McpPool::from_config_path(&path).unwrap();
3758 let drops = Arc::new(AtomicUsize::new(0));
3759 let mut conn = test_connection(Box::new(DropCountingTransport {
3760 drops: Arc::clone(&drops),
3761 }));
3762 conn.name = "local".to_string();
3763 conn.config = pool.config.servers.get("local").unwrap().clone();
3764 pool.connections.insert("local".to_string(), conn);
3765 let mut runtime_config = test_server_config();
3766 runtime_config.command = Some("runtime-server".to_string());
3767 pool.add_runtime_server_config("runtime".to_string(), runtime_config)
3768 .unwrap();
3769 let generation_before = pool.catalog_generation.load(AtomicOrdering::SeqCst);
3770
3771 pool.force_reload_config_sources().unwrap();
3772 let errors = pool.connect_all().await;
3773
3774 assert!(
3775 errors.is_empty(),
3776 "disabled config should not connect: {errors:?}"
3777 );
3778 assert_eq!(drops.load(AtomicOrdering::SeqCst), 1);
3779 assert!(!pool.connections.contains_key("local"));
3780 assert!(pool.server_names().contains(&"runtime".to_string()));
3781 assert_eq!(
3782 pool.catalog_generation.load(AtomicOrdering::SeqCst),
3783 generation_before + 1,
3784 "explicit reload must invalidate every previously advertised route"
3785 );
3786 }
3787
3788 #[tokio::test]
3789 async fn config_source_switch_preserves_dynamic_servers_in_the_shared_pool() {
3790 let dir = tempfile::tempdir().unwrap();
3791 let workspace = dir.path().join("workspace");
3792 std::fs::create_dir_all(&workspace).unwrap();
3793 let initial_path = dir.path().join("initial.json");
3794 let invalid_path = dir.path().join("invalid.json");
3795 let replacement_path = dir.path().join("replacement.json");
3796 std::fs::write(
3797 &initial_path,
3798 r#"{"servers":{"local":{"command":"node","disabled":true}}}"#,
3799 )
3800 .unwrap();
3801 std::fs::write(&invalid_path, r#"{"servers":{"broken": trailing}}"#).unwrap();
3802 std::fs::write(&replacement_path, r#"{"servers":{}}"#).unwrap();
3803 let plugins = Arc::new(crate::plugins::PluginRegistry::empty(&workspace));
3804 let mut pool = McpPool::from_config_path_with_workspace_and_plugins(
3805 &initial_path,
3806 &workspace,
3807 Arc::clone(&plugins),
3808 )
3809 .unwrap();
3810 let mut runtime_config = test_server_config();
3811 runtime_config.command = Some("runtime-server".to_string());
3812 pool.add_runtime_server_config("runtime".to_string(), runtime_config)
3813 .unwrap();
3814 let drops = Arc::new(AtomicUsize::new(0));
3815 let mut conn = test_connection(Box::new(DropCountingTransport {
3816 drops: Arc::clone(&drops),
3817 }));
3818 conn.name = "local".to_string();
3819 conn.config = pool.config.servers.get("local").unwrap().clone();
3820 pool.connections.insert("local".to_string(), conn);
3821 let generation_before = pool.catalog_generation.load(AtomicOrdering::SeqCst);
3822
3823 pool.switch_workspace_config_source(&invalid_path, &workspace, Arc::clone(&plugins))
3824 .expect_err("malformed replacement must fail closed");
3825 assert_eq!(pool.config_sources.first(), Some(&initial_path));
3826 assert!(pool.connections.contains_key("local"));
3827 assert_eq!(drops.load(AtomicOrdering::SeqCst), 0);
3828 assert_eq!(
3829 pool.catalog_generation.load(AtomicOrdering::SeqCst),
3830 generation_before
3831 );
3832
3833 pool.switch_workspace_config_source(&replacement_path, &workspace, plugins)
3834 .unwrap();
3835 let errors = pool.connect_all().await;
3836
3837 assert!(errors.is_empty());
3838 assert!(pool.server_names().contains(&"runtime".to_string()));
3839 assert_eq!(pool.config_sources.first(), Some(&replacement_path));
3840 assert_eq!(drops.load(AtomicOrdering::SeqCst), 1);
3841 }
3842
3843 /// #1267 part 2: hash-based comparison must be stable for byte-identical
3844 /// configs and distinct for differing configs.
3845 #[test]
3846 fn hash_mcp_config_is_stable_and_change_sensitive() {
3847 let a = McpConfig::default();
3848 let b = McpConfig::default();
3849 assert_eq!(hash_mcp_config(&a), hash_mcp_config(&b));
3850 let mut c = McpConfig::default();
3851 c.servers.insert(
3852 "x".into(),
3853 McpServerConfig {
3854 command: Some("/bin/echo".into()),
3855 args: vec!["hi".into()],
3856 env: Default::default(),
3857 cwd: None,
3858 url: None,
3859 transport: None,
3860 connect_timeout: None,
3861 execute_timeout: None,
3862 read_timeout: None,
3863 disabled: false,
3864 enabled: true,
3865 required: false,
3866 enabled_tools: Vec::new(),
3867 disabled_tools: Vec::new(),
3868 headers: HashMap::new(),
3869 env_headers: HashMap::new(),
3870 bearer_token_env_var: None,
3871 scopes: Vec::new(),
3872 oauth: None,
3873 oauth_resource: None,
3874 reviewed_plugin: None,
3875 runtime_added: false,
3876 allow_private_network: false,
3877 },
3878 );
3879 assert_ne!(
3880 hash_mcp_config(&a),
3881 hash_mcp_config(&c),
3882 "hash must change when servers map changes"
3883 );
3884 }
3885
3886 /// #1267 part 2: `hash_mcp_config` is the *only* thing standing between a
3887 /// touched-but-unchanged config file and a full teardown of every live MCP
3888 /// connection (stdio children included). `McpConfig::servers`, `env`,
3889 /// `headers`, and `env_headers` are all `HashMap`s, and two `HashMap`s built
3890 /// separately in one process iterate in different orders — so hashing the
3891 /// config's `serde_json` bytes straight from the struct is not
3892 /// content-addressed once a map holds more than one entry. Parse the same
3893 /// bytes repeatedly and require one hash.
3894 #[test]
3895 fn hash_mcp_config_is_order_independent_across_identical_parses() {
3896 let raw = r#"{
3897 "servers": {
3898 "alpha": { "command": "a", "env": { "A": "1", "B": "2", "C": "3", "D": "4" } },
3899 "bravo": { "command": "b" },
3900 "charlie": { "command": "c" },
3901 "delta": { "command": "d" },
3902 "echo": { "command": "e" },
3903 "foxtrot": { "command": "f" },
3904 "golf": { "command": "g" },
3905 "hotel": { "command": "h" },
3906 "india": { "command": "i" },
3907 "juliett": { "command": "j" }
3908 }
3909 }"#;
3910 let hashes: HashSet<u64> = (0..16)
3911 .map(|_| {
3912 let parsed: McpConfig = serde_json::from_str(raw).expect("fixture parses");
3913 hash_mcp_config(&parsed)
3914 })
3915 .collect();
3916 assert_eq!(
3917 hashes.len(),
3918 1,
3919 "byte-identical MCP config must hash identically; got {} distinct hashes",
3920 hashes.len()
3921 );
3922 }
3923
3924 /// The same invariant for the nested per-server maps: an unchanged server
3925 /// whose `env` / `headers` hold several entries must not look changed.
3926 #[test]
3927 fn hash_mcp_config_is_order_independent_for_nested_server_maps() {
3928 let raw = r#"{
3929 "servers": {
3930 "only": {
3931 "url": "https://example.invalid/mcp",
3932 "headers": { "H1": "1", "H2": "2", "H3": "3", "H4": "4", "H5": "5", "H6": "6" },
3933 "env_headers": { "E1": "V1", "E2": "V2", "E3": "V3", "E4": "V4" }
3934 }
3935 }
3936 }"#;
3937 let hashes: HashSet<u64> = (0..16)
3938 .map(|_| {
3939 let parsed: McpConfig = serde_json::from_str(raw).expect("fixture parses");
3940 hash_mcp_config(&parsed)
3941 })
3942 .collect();
3943 assert_eq!(
3944 hashes.len(),
3945 1,
3946 "byte-identical MCP config must hash identically; got {} distinct hashes",
3947 hashes.len()
3948 );
3949 }
3950
3951 /// #1319: discovered tools must be sorted by name so the prompt prefix
3952 /// is stable across runs (cache-hit stability), even when the server
3953 /// returns them in arbitrary or paginated order.
3954 #[tokio::test]
3955 async fn discover_tools_sorts_by_name_for_cache_stability() {
3956 let sent = Arc::new(Mutex::new(Vec::new()));
3957 let transport = ScriptedValueTransport {
3958 sent: Arc::clone(&sent),
3959 responses: VecDeque::from([
3960 json_frame(serde_json::json!({
3961 "jsonrpc": "2.0",
3962 "id": 1,
3963 "result": {
3964 "tools": [
3965 { "name": "zeta", "inputSchema": {} },
3966 { "name": "alpha", "inputSchema": {} }
3967 ],
3968 "nextCursor": "page-2"
3969 }
3970 })),
3971 json_frame(serde_json::json!({
3972 "jsonrpc": "2.0",
3973 "id": 2,
3974 "result": {
3975 "tools": [
3976 { "name": "mu", "inputSchema": {} },
3977 { "name": "beta", "inputSchema": {} }
3978 ]
3979 }
3980 })),
3981 ]),
3982 };
3983 let mut conn = test_connection(Box::new(transport));
3984 conn.tools = conn
3985 .discover_tools(&mut McpCatalogBudget::new())
3986 .await
3987 .expect("discover");
3988
3989 let names: Vec<&str> = conn.tools.iter().map(|t| t.name.as_str()).collect();
3990 assert_eq!(
3991 names,
3992 vec!["alpha", "beta", "mu", "zeta"],
3993 "tools must be sorted by name regardless of server order or pagination"
3994 );
3995 }
3996
3997 #[tokio::test]
3998 async fn discover_tools_rejects_a_repeated_pagination_cursor_without_publishing_partials() {
3999 let transport = ScriptedValueTransport {
4000 sent: Arc::new(Mutex::new(Vec::new())),
4001 responses: VecDeque::from([
4002 json_frame(serde_json::json!({
4003 "jsonrpc": "2.0",
4004 "id": 1,
4005 "result": {
4006 "tools": [{ "name": "first", "inputSchema": {} }],
4007 "nextCursor": "same"
4008 }
4009 })),
4010 json_frame(serde_json::json!({
4011 "jsonrpc": "2.0",
4012 "id": 2,
4013 "result": {
4014 "tools": [{ "name": "second", "inputSchema": {} }],
4015 "nextCursor": "same"
4016 }
4017 })),
4018 ]),
4019 };
4020 let mut conn = test_connection(Box::new(transport));
4021
4022 let error = conn
4023 .discover_tools(&mut McpCatalogBudget::new())
4024 .await
4025 .expect_err("repeated cursor must abort discovery");
4026 assert!(error.to_string().contains("repeated pagination cursor"));
4027 assert!(
4028 conn.tools.is_empty(),
4029 "an aborted catalogue must not publish attacker-controlled partial entries"
4030 );
4031 }
4032
4033 #[test]
4034 fn mcp_tool_description_formatter_is_one_line_and_unicode_safe() {
4035 let long_cjk = format!("{}\n这行不应显示", "鲸".repeat(81));
4036 assert_eq!(
4037 format_mcp_tool_description(Some(&long_cjk)),
4038 format!(": {}...", "鲸".repeat(80))
4039 );
4040 assert_eq!(
4041 format_mcp_tool_description(Some("第一行\r\n第二行")),
4042 ": 第一行"
4043 );
4044 assert_eq!(format_mcp_tool_description(Some(" \nignored")), "");
4045 assert_eq!(format_mcp_tool_description(None), "");
4046 }
4047
4048 #[tokio::test]
4049 async fn discover_all_honors_tools_only_server_capabilities() {
4050 let sent = Arc::new(Mutex::new(Vec::new()));
4051 let transport = ScriptedValueTransport {
4052 sent: Arc::clone(&sent),
4053 responses: VecDeque::from([
4054 json_frame(serde_json::json!({
4055 "jsonrpc": "2.0",
4056 "id": 1,
4057 "result": {
4058 "protocolVersion": "2024-11-05",
4059 "serverInfo": {"name": "tools-only", "version": "1.0.0"},
4060 "capabilities": {"tools": {}}
4061 }
4062 })),
4063 json_frame(serde_json::json!({
4064 "jsonrpc": "2.0",
4065 "id": 2,
4066 "result": {
4067 "tools": [{"name": "idea_search", "inputSchema": {}}]
4068 }
4069 })),
4070 ]),
4071 };
4072 let mut conn = test_connection(Box::new(transport));
4073
4074 conn.initialize().await.expect("initialize");
4075 conn.discover_all().await.expect("discover tools");
4076
4077 assert_eq!(
4078 conn.server_capabilities,
4079 Some(McpServerCapabilities {
4080 tools: true,
4081 resources: false,
4082 prompts: false,
4083 })
4084 );
4085 assert_eq!(conn.tools.len(), 1);
4086 assert!(conn.resources.is_empty());
4087 assert!(conn.resource_templates.is_empty());
4088 assert!(conn.prompts.is_empty());
4089 let methods: Vec<_> = sent
4090 .lock()
4091 .unwrap()
4092 .iter()
4093 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
4094 .map(str::to_string)
4095 .collect();
4096 assert_eq!(
4097 methods,
4098 ["initialize", "notifications/initialized", "tools/list"]
4099 );
4100 }
4101
4102 #[tokio::test]
4103 async fn discover_all_populates_every_advertised_capability() {
4104 let sent = Arc::new(Mutex::new(Vec::new()));
4105 let transport = ScriptedValueTransport {
4106 sent: Arc::clone(&sent),
4107 responses: VecDeque::from([
4108 json_frame(serde_json::json!({
4109 "jsonrpc": "2.0",
4110 "id": 1,
4111 "result": {
4112 "protocolVersion": "2024-11-05",
4113 "serverInfo": {"name": "full", "version": "1.0.0"},
4114 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
4115 }
4116 })),
4117 json_frame(serde_json::json!({
4118 "jsonrpc": "2.0",
4119 "id": 2,
4120 "result": {"tools": [{"name": "search", "inputSchema": {}}]}
4121 })),
4122 json_frame(serde_json::json!({
4123 "jsonrpc": "2.0",
4124 "id": 3,
4125 "result": {"resources": [{"uri": "file:///readme", "name": "readme"}]}
4126 })),
4127 json_frame(serde_json::json!({
4128 "jsonrpc": "2.0",
4129 "id": 4,
4130 "result": {
4131 "resourceTemplates": [{"uriTemplate": "file:///{path}", "name": "file"}]
4132 }
4133 })),
4134 json_frame(serde_json::json!({
4135 "jsonrpc": "2.0",
4136 "id": 5,
4137 "result": {"prompts": [{"name": "review"}]}
4138 })),
4139 ]),
4140 };
4141 let mut conn = test_connection(Box::new(transport));
4142
4143 conn.initialize().await.expect("initialize");
4144 conn.discover_all().await.expect("discover all");
4145
4146 assert_eq!(
4147 conn.server_capabilities,
4148 Some(McpServerCapabilities {
4149 tools: true,
4150 resources: true,
4151 prompts: true,
4152 })
4153 );
4154 assert_eq!(conn.tools.len(), 1);
4155 assert_eq!(conn.resources.len(), 1);
4156 assert_eq!(conn.resource_templates.len(), 1);
4157 assert_eq!(conn.prompts.len(), 1);
4158 let methods: Vec<_> = sent
4159 .lock()
4160 .unwrap()
4161 .iter()
4162 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
4163 .map(str::to_string)
4164 .collect();
4165 assert_eq!(
4166 methods,
4167 [
4168 "initialize",
4169 "notifications/initialized",
4170 "tools/list",
4171 "resources/list",
4172 "resources/templates/list",
4173 "prompts/list",
4174 ]
4175 );
4176 }
4177
4178 #[tokio::test]
4179 async fn legacy_optional_discovery_hangs_are_bounded_and_fail_soft() {
4180 let sent = Arc::new(Mutex::new(Vec::new()));
4181 let transport = ScriptedThenHangingTransport {
4182 sent: Arc::clone(&sent),
4183 responses: VecDeque::from([json_frame(serde_json::json!({
4184 "jsonrpc": "2.0",
4185 "id": 1,
4186 "result": {"tools": [{"name": "search", "inputSchema": {}}]}
4187 }))]),
4188 };
4189 let mut conn = test_connection(Box::new(transport));
4190 conn.discovery_timeout = Duration::from_millis(60);
4191
4192 let started = tokio::time::Instant::now();
4193 conn.discover_all()
4194 .await
4195 .expect("hung optional methods must not fail discovery");
4196
4197 assert_eq!(conn.server_capabilities, None);
4198 assert_eq!(conn.tools.len(), 1);
4199 assert!(
4200 started.elapsed() < Duration::from_secs(1),
4201 "optional discovery exceeded its bounded budget: {:?}",
4202 started.elapsed()
4203 );
4204 let methods: Vec<_> = sent
4205 .lock()
4206 .unwrap()
4207 .iter()
4208 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
4209 .map(str::to_string)
4210 .collect();
4211 assert_eq!(
4212 methods,
4213 [
4214 "tools/list",
4215 "resources/list",
4216 "resources/templates/list",
4217 "prompts/list",
4218 ]
4219 );
4220 }
4221
4222 #[test]
4223 fn manager_snapshot_preserves_advertised_and_legacy_capability_provenance() {
4224 let advertised_config = test_server_config();
4225 let legacy_config = test_server_config();
4226 let config = McpConfig {
4227 servers: HashMap::from([
4228 ("advertised".to_string(), advertised_config.clone()),
4229 ("legacy".to_string(), legacy_config.clone()),
4230 ]),
4231 ..McpConfig::default()
4232 };
4233 let mut pool = McpPool::new(config.clone());
4234 let drops = Arc::new(AtomicUsize::new(0));
4235
4236 let mut advertised = test_connection(Box::new(DropCountingTransport {
4237 drops: Arc::clone(&drops),
4238 }));
4239 advertised.name = "advertised".to_string();
4240 advertised.config = advertised_config;
4241 advertised.server_capabilities = Some(McpServerCapabilities {
4242 tools: true,
4243 resources: false,
4244 prompts: true,
4245 });
4246 pool.connections
4247 .insert("advertised".to_string(), advertised);
4248
4249 let mut legacy = test_connection(Box::new(DropCountingTransport { drops }));
4250 legacy.name = "legacy".to_string();
4251 legacy.config = legacy_config;
4252 pool.connections.insert("legacy".to_string(), legacy);
4253
4254 let errors = HashMap::new();
4255 let snapshot = snapshot_from_config(
4256 Path::new("mcp.json"),
4257 true,
4258 false,
4259 &config,
4260 Some((&pool, &errors)),
4261 );
4262
4263 assert_eq!(
4264 snapshot.servers[0].capability_metadata,
4265 McpServerCapabilityMetadata::Advertised(McpServerCapabilities {
4266 tools: true,
4267 resources: false,
4268 prompts: true,
4269 })
4270 );
4271 assert_eq!(
4272 snapshot.servers[1].capability_metadata,
4273 McpServerCapabilityMetadata::LegacyFallback
4274 );
4275 }
4276
4277 #[tokio::test]
4278 async fn mcp_pool_call_tool_preserves_tool_names_with_dashes() {
4279 let sent = Arc::new(Mutex::new(Vec::new()));
4280 let transport = ScriptedValueTransport {
4281 sent: Arc::clone(&sent),
4282 responses: VecDeque::from([json_frame(serde_json::json!({
4283 "jsonrpc": "2.0",
4284 "id": 1,
4285 "result": {"ok": true}
4286 }))]),
4287 };
4288 let mut conn = test_connection(Box::new(transport));
4289 conn.name = "dephy".to_string();
4290 conn.tools = vec![McpTool {
4291 name: "company--search".to_string(),
4292 description: None,
4293 input_schema: serde_json::json!({}),
4294 annotations: None,
4295 }];
4296
4297 let mut pool = McpPool::new(McpConfig {
4298 timeouts: McpTimeouts::default(),
4299 servers: HashMap::new(),
4300 });
4301 pool.connections.insert("dephy".to_string(), conn);
4302
4303 let result = pool
4304 .call_tool(
4305 "mcp_dephy_company--search",
4306 serde_json::json!({"query": "dephy"}),
4307 )
4308 .await
4309 .unwrap();
4310
4311 assert_eq!(result, serde_json::json!({"ok": true}));
4312 let sent = sent.lock().unwrap();
4313 assert_eq!(sent[0]["method"], "tools/call");
4314 assert_eq!(sent[0]["params"]["name"], "company--search");
4315 assert_eq!(
4316 sent[0]["params"]["arguments"],
4317 serde_json::json!({"query": "dephy"})
4318 );
4319 }
4320
4321 #[tokio::test]
4322 async fn mcp_pool_rejects_unadvertised_tool_without_sending_tools_call() {
4323 let sent = Arc::new(Mutex::new(Vec::new()));
4324 let transport = ScriptedValueTransport {
4325 sent: Arc::clone(&sent),
4326 // A malicious server could implement this hidden method, but local
4327 // catalog authorization must prevent the transport from seeing it.
4328 responses: VecDeque::from([json_frame(serde_json::json!({
4329 "jsonrpc": "2.0", "id": 1, "result": {"deleted": true}
4330 }))]),
4331 };
4332 let mut conn = test_connection(Box::new(transport));
4333 conn.name = "spy".to_string();
4334 conn.tools = vec![McpTool {
4335 name: "read".to_string(),
4336 description: None,
4337 input_schema: serde_json::json!({}),
4338 annotations: None,
4339 }];
4340 let mut pool = McpPool::new(McpConfig::default());
4341 pool.connections.insert("spy".to_string(), conn);
4342
4343 let error = pool
4344 .call_tool("mcp_spy_delete", serde_json::json!({}))
4345 .await
4346 .expect_err("unadvertised hidden tool must fail locally");
4347 assert!(error.to_string().contains("Unknown MCP tool name"));
4348 assert!(sent.lock().unwrap().is_empty(), "zero tools/call requests");
4349 }
4350
4351 #[tokio::test]
4352 async fn mcp_pool_binds_prompts_and_resources_to_advertised_catalog() {
4353 let sent = Arc::new(Mutex::new(Vec::new()));
4354 let transport = ScriptedValueTransport {
4355 sent: Arc::clone(&sent),
4356 responses: VecDeque::from([json_frame(serde_json::json!({
4357 "jsonrpc": "2.0", "id": 1, "result": {"contents": []}
4358 }))]),
4359 };
4360 let mut conn = test_connection(Box::new(transport));
4361 conn.name = "catalog".to_string();
4362 conn.prompts = vec![McpPrompt {
4363 name: "review".to_string(),
4364 description: None,
4365 arguments: Vec::new(),
4366 }];
4367 conn.resources = vec![McpResource {
4368 uri: "file:///readme".to_string(),
4369 name: "readme".to_string(),
4370 description: None,
4371 mime_type: None,
4372 }];
4373 conn.resource_templates = vec![McpResourceTemplate {
4374 uri_template: "repo://item/{id}".to_string(),
4375 name: "item".to_string(),
4376 description: None,
4377 mime_type: None,
4378 }];
4379 let mut pool = McpPool::new(McpConfig::default());
4380 pool.connections.insert("catalog".to_string(), conn);
4381
4382 pool.get_prompt("catalog", "hidden", serde_json::json!({}))
4383 .await
4384 .expect_err("hidden prompt must fail locally");
4385 pool.read_resource("catalog", "file:///hidden")
4386 .await
4387 .expect_err("hidden literal resource must fail locally");
4388 assert!(sent.lock().unwrap().is_empty());
4389
4390 let result = pool
4391 .read_resource("catalog", "repo://item/42")
4392 .await
4393 .expect("exact advertised template expansion is callable");
4394 assert_eq!(result, serde_json::json!({"contents": []}));
4395 let sent = sent.lock().unwrap();
4396 assert_eq!(sent.len(), 1);
4397 assert_eq!(sent[0]["method"], "resources/read");
4398 assert_eq!(sent[0]["params"]["uri"], "repo://item/42");
4399 }
4400
4401 #[tokio::test]
4402 async fn mcp_pool_call_tool_preserves_server_names_with_underscores() {
4403 let sent = Arc::new(Mutex::new(Vec::new()));
4404 let transport = ScriptedValueTransport {
4405 sent: Arc::clone(&sent),
4406 responses: VecDeque::from([json_frame(serde_json::json!({
4407 "jsonrpc": "2.0",
4408 "id": 1,
4409 "result": {"ok": true}
4410 }))]),
4411 };
4412 let mut conn = test_connection(Box::new(transport));
4413 conn.name = "my_db".to_string();
4414 conn.tools = vec![McpTool {
4415 name: "execute_sql".to_string(),
4416 description: None,
4417 input_schema: serde_json::json!({}),
4418 annotations: None,
4419 }];
4420
4421 let mut pool = McpPool::new(McpConfig {
4422 timeouts: McpTimeouts::default(),
4423 servers: HashMap::new(),
4424 });
4425 pool.connections.insert("my_db".to_string(), conn);
4426
4427 let result = pool
4428 .call_tool(
4429 "mcp_my_db_execute_sql",
4430 serde_json::json!({"query": "select 1"}),
4431 )
4432 .await
4433 .unwrap();
4434
4435 assert_eq!(result, serde_json::json!({"ok": true}));
4436 let sent = sent.lock().unwrap();
4437 assert_eq!(sent[0]["method"], "tools/call");
4438 assert_eq!(sent[0]["params"]["name"], "execute_sql");
4439 assert_eq!(
4440 sent[0]["params"]["arguments"],
4441 serde_json::json!({"query": "select 1"})
4442 );
4443 }
4444
4445 #[tokio::test]
4446 async fn mcp_pool_hides_and_rejects_ambiguous_model_tool_names() {
4447 let sent_short = Arc::new(Mutex::new(Vec::new()));
4448 let short_transport = ScriptedValueTransport {
4449 sent: Arc::clone(&sent_short),
4450 responses: VecDeque::from([json_frame(serde_json::json!({
4451 "jsonrpc": "2.0",
4452 "id": 1,
4453 "result": {"short": true}
4454 }))]),
4455 };
4456 let mut short_conn = test_connection(Box::new(short_transport));
4457 short_conn.name = "my".to_string();
4458 short_conn.tools = vec![McpTool {
4459 name: "db_execute_sql".to_string(),
4460 description: None,
4461 input_schema: serde_json::json!({}),
4462 annotations: None,
4463 }];
4464
4465 let sent_long = Arc::new(Mutex::new(Vec::new()));
4466 let long_transport = ScriptedValueTransport {
4467 sent: Arc::clone(&sent_long),
4468 responses: VecDeque::from([json_frame(serde_json::json!({
4469 "jsonrpc": "2.0",
4470 "id": 1,
4471 "result": {"long": true}
4472 }))]),
4473 };
4474 let mut long_conn = test_connection(Box::new(long_transport));
4475 long_conn.name = "my_db".to_string();
4476 long_conn.tools = vec![McpTool {
4477 name: "execute_sql".to_string(),
4478 description: None,
4479 input_schema: serde_json::json!({}),
4480 annotations: None,
4481 }];
4482
4483 let mut pool = McpPool::new(McpConfig {
4484 timeouts: McpTimeouts::default(),
4485 servers: HashMap::new(),
4486 });
4487 pool.connections.insert("my".to_string(), short_conn);
4488 pool.connections.insert("my_db".to_string(), long_conn);
4489
4490 assert!(
4491 pool.all_tools().is_empty(),
4492 "ambiguous names must never be advertised to the model"
4493 );
4494 let error = pool
4495 .call_tool(
4496 "mcp_my_db_execute_sql",
4497 serde_json::json!({"query": "select 1"}),
4498 )
4499 .await
4500 .expect_err("ambiguous tool route must fail closed");
4501
4502 assert!(error.to_string().contains("Ambiguous MCP tool name"));
4503 assert!(
4504 sent_short.lock().unwrap().is_empty(),
4505 "neither authority may receive an ambiguous tool call"
4506 );
4507 assert!(
4508 sent_long.lock().unwrap().is_empty(),
4509 "neither authority may receive an ambiguous tool call"
4510 );
4511 }
4512
4513 #[tokio::test]
4514 async fn json_rpc_session_error_is_marked_stale() {
4515 let sent = Arc::new(Mutex::new(Vec::new()));
4516 let transport = ScriptedValueTransport {
4517 sent: Arc::clone(&sent),
4518 responses: VecDeque::from([json_frame(serde_json::json!({
4519 "jsonrpc": "2.0",
4520 "id": 1,
4521 "error": {
4522 "code": -32001,
4523 "message": "MCP session expired"
4524 }
4525 }))]),
4526 };
4527 let mut conn = test_connection(Box::new(transport));
4528
4529 let err = conn
4530 .call_tool("search", serde_json::json!({"query": "dephy"}), 1)
4531 .await
4532 .expect_err("session error should fail");
4533
4534 assert!(
4535 is_mcp_stale_session_error(&err),
4536 "JSON-RPC session error should be retryable, got: {err:#}"
4537 );
4538 }
4539
4540 #[test]
4541 fn sse_transport_closed_is_retryable() {
4542 let err = anyhow::anyhow!("SSE transport closed");
4543 assert!(
4544 is_mcp_stale_session_error(&err),
4545 "closed SSE stream should force reconnect before retry"
4546 );
4547 }
4548
4549 #[test]
4550 fn stdio_transport_closed_is_retryable() {
4551 let err = anyhow::anyhow!("Stdio transport closed (exit status: 1)");
4552 assert!(
4553 is_mcp_stale_session_error(&err),
4554 "dead stdio child should force reconnect before retry"
4555 );
4556 }
4557
4558 #[test]
4559 fn legacy_sse_post_disconnect_is_retryable() {
4560 let err = anyhow::anyhow!(
4561 "MCP SSE POST send failed (transport=sse endpoint=http://127.0.0.1:123/messages): connection closed before message completed"
4562 );
4563 assert!(
4564 is_mcp_stale_session_error(&err),
4565 "closed legacy SSE POST should force reconnect before retry"
4566 );
4567
4568 let err = anyhow::anyhow!(
4569 "MCP SSE POST send failed (transport=sse endpoint=http://127.0.0.1:123/messages): connection reset by peer"
4570 );
4571 assert!(
4572 is_mcp_stale_session_error(&err),
4573 "reset legacy SSE POST should force reconnect before retry"
4574 );
4575
4576 let err = anyhow::anyhow!(
4577 "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."
4578 );
4579 assert!(
4580 is_mcp_stale_session_error(&err),
4581 "Windows reset wording should force reconnect before retry"
4582 );
4583 }
4584
4585 #[tokio::test]
4586 async fn discover_all_ignores_unsupported_optional_capabilities() {
4587 let sent = Arc::new(Mutex::new(Vec::new()));
4588 let transport = ScriptedValueTransport {
4589 sent: Arc::clone(&sent),
4590 responses: VecDeque::from([
4591 json_frame(serde_json::json!({
4592 "jsonrpc": "2.0",
4593 "id": 1,
4594 "result": {
4595 "tools": [
4596 { "name": "search", "inputSchema": {} }
4597 ]
4598 }
4599 })),
4600 json_frame(serde_json::json!({
4601 "jsonrpc": "2.0",
4602 "id": 2,
4603 "error": {
4604 "code": -32601,
4605 "message": "resources not supported"
4606 }
4607 })),
4608 json_frame(serde_json::json!({
4609 "jsonrpc": "2.0",
4610 "id": 3,
4611 "error": {
4612 "code": -32601,
4613 "message": "resource templates not supported"
4614 }
4615 })),
4616 json_frame(serde_json::json!({
4617 "jsonrpc": "2.0",
4618 "id": 4,
4619 "error": {
4620 "code": -32601,
4621 "message": "prompts not supported"
4622 }
4623 })),
4624 ]),
4625 };
4626 let mut conn = test_connection(Box::new(transport));
4627 conn.server_capabilities = Some(McpServerCapabilities {
4628 tools: true,
4629 resources: true,
4630 prompts: true,
4631 });
4632
4633 conn.discover_all().await.expect("discover");
4634
4635 assert_eq!(conn.tools.len(), 1);
4636 assert_eq!(conn.tools[0].name, "search");
4637 assert!(conn.resources.is_empty());
4638 assert!(conn.resource_templates.is_empty());
4639 assert!(conn.prompts.is_empty());
4640 let methods: Vec<_> = sent
4641 .lock()
4642 .unwrap()
4643 .iter()
4644 .filter_map(|message| message.get("method").and_then(|method| method.as_str()))
4645 .map(str::to_string)
4646 .collect();
4647 assert_eq!(
4648 methods,
4649 [
4650 "tools/list",
4651 "resources/list",
4652 "resources/templates/list",
4653 "prompts/list",
4654 ]
4655 );
4656 }
4657
4658 #[tokio::test]
4659 async fn discover_all_keeps_advertised_tool_discovery_required() {
4660 let sent = Arc::new(Mutex::new(Vec::new()));
4661 let transport = ScriptedValueTransport {
4662 sent,
4663 responses: VecDeque::from([json_frame(serde_json::json!({
4664 "jsonrpc": "2.0",
4665 "id": 1,
4666 "error": {"code": -32601, "message": "tools not supported"}
4667 }))]),
4668 };
4669 let mut conn = test_connection(Box::new(transport));
4670 conn.server_capabilities = Some(McpServerCapabilities {
4671 tools: true,
4672 resources: false,
4673 prompts: false,
4674 });
4675
4676 let error = conn
4677 .discover_all()
4678 .await
4679 .expect_err("advertised tools/list failure must fail discovery");
4680
4681 assert!(
4682 error.to_string().contains("MCP error in 'tools/list'"),
4683 "unexpected error: {error:#}"
4684 );
4685 }
4686
4687 /// #1244: when an MCP stdio server fails to spawn, the underlying OS
4688 /// error (e.g. ENOENT for a missing binary) must reach the user via the
4689 /// snapshot.error string. Regression test for `err.to_string()` dropping
4690 /// the anyhow chain — without `{err:#}` the user sees only the opaque
4691 /// wrapper "MCP stdio spawn failed (...)" and has nothing to act on.
4692 #[tokio::test]
4693 async fn discover_snapshot_includes_underlying_spawn_error_in_chain() {
4694 let dir = tempfile::tempdir().unwrap();
4695 let path = dir.path().join("mcp.json");
4696 fs::write(
4697 &path,
4698 r#"{
4699 "mcpServers": {
4700 "broken": {
4701 "command": "codewhale-tui-test-this-binary-does-not-exist-9f8e7d6c5b4a",
4702 "args": []
4703 }
4704 }
4705 }"#,
4706 )
4707 .unwrap();
4708
4709 let snapshot = discover_manager_snapshot(&path, None, false).await.unwrap();
4710 let server = snapshot
4711 .servers
4712 .iter()
4713 .find(|s| s.name == "broken")
4714 .expect("broken server should appear in snapshot");
4715 let err = server
4716 .error
4717 .as_deref()
4718 .expect("broken server should have an error");
4719 let lowered = err.to_lowercase();
4720 assert!(
4721 lowered.contains("os error")
4722 || lowered.contains("not found")
4723 || lowered.contains("no such"),
4724 "expected underlying spawn error in chain, got: {err}"
4725 );
4726 }
4727
4728 #[tokio::test]
4729 async fn discover_snapshot_explains_a_missing_node_runtime() {
4730 let dir = tempfile::tempdir().unwrap();
4731 let config_path = dir.path().join("mcp.json");
4732 let missing_node = dir
4733 .path()
4734 .join(if cfg!(windows) { "node.exe" } else { "node" });
4735 fs::write(
4736 &config_path,
4737 serde_json::to_vec(&serde_json::json!({
4738 "mcpServers": { "computer": { "command": missing_node, "args": [] } }
4739 }))
4740 .unwrap(),
4741 )
4742 .unwrap();
4743 let snapshot = discover_manager_snapshot(&config_path, None, false)
4744 .await
4745 .unwrap();
4746 let error = snapshot
4747 .servers
4748 .iter()
4749 .find(|server| server.name == "computer")
4750 .unwrap()
4751 .error
4752 .as_ref()
4753 .unwrap();
4754 assert!(error.contains("Node.js 20 or newer"), "{error}");
4755 assert!(error.contains("https://nodejs.org/"), "{error}");
4756 }
4757
4758 /// The same guarantee for a server the user marked `required`. `connect_all`
4759 /// appends a generic "required MCP server failed to initialize" entry after
4760 /// the real per-server connect error, and every snapshot path folds the
4761 /// returned pairs into a `HashMap<name, message>` — so the later, contentless
4762 /// entry overwrites the diagnosis. Marking a server required must not blind
4763 /// the user to *why* it did not start.
4764 #[tokio::test]
4765 async fn required_server_snapshot_keeps_the_real_spawn_error() {
4766 let dir = tempfile::tempdir().unwrap();
4767 let path = dir.path().join("mcp.json");
4768 fs::write(
4769 &path,
4770 r#"{
4771 "mcpServers": {
4772 "broken": {
4773 "command": "codewhale-tui-test-this-binary-does-not-exist-9f8e7d6c5b4a",
4774 "args": [],
4775 "required": true
4776 }
4777 }
4778 }"#,
4779 )
4780 .unwrap();
4781
4782 let snapshot = discover_manager_snapshot(&path, None, false).await.unwrap();
4783 let server = snapshot
4784 .servers
4785 .iter()
4786 .find(|s| s.name == "broken")
4787 .expect("broken server should appear in snapshot");
4788 let err = server
4789 .error
4790 .as_deref()
4791 .expect("broken server should have an error");
4792 let lowered = err.to_lowercase();
4793 assert!(
4794 lowered.contains("os error")
4795 || lowered.contains("not found")
4796 || lowered.contains("no such"),
4797 "required server must still report why it failed, got: {err}"
4798 );
4799 }
4800
4801 /// `connect_all` must report one error per failed server. A `required`
4802 /// server that already failed to connect got a second, contentless entry
4803 /// appended for the same name.
4804 #[tokio::test]
4805 async fn connect_all_reports_one_error_per_failed_required_server() {
4806 let dir = tempfile::tempdir().unwrap();
4807 let path = dir.path().join("mcp.json");
4808 fs::write(
4809 &path,
4810 r#"{
4811 "mcpServers": {
4812 "broken": {
4813 "command": "codewhale-tui-test-this-binary-does-not-exist-9f8e7d6c5b4a",
4814 "args": [],
4815 "required": true
4816 }
4817 }
4818 }"#,
4819 )
4820 .unwrap();
4821
4822 let mut pool = McpPool::from_config_path(&path).unwrap();
4823 let errors = pool.connect_all().await;
4824 let for_broken: Vec<_> = errors.iter().filter(|(name, _)| name == "broken").collect();
4825 assert_eq!(
4826 for_broken.len(),
4827 1,
4828 "expected exactly one error for 'broken', got: {:?}",
4829 errors
4830 .iter()
4831 .map(|(name, err)| format!("{name}: {err:#}"))
4832 .collect::<Vec<_>>()
4833 );
4834 let rendered = format!("{:#}", for_broken[0].1).to_lowercase();
4835 assert!(
4836 rendered.contains("spawn failed"),
4837 "the single error must be the real cause, got: {rendered}"
4838 );
4839 }
4840
4841 /// A dead server must not be re-dialed on every turn.
4842 ///
4843 /// The turn loop rebuilds the tool catalog on each user message, and that
4844 /// path calls `connect_all`. Before the cooldown, a wall of unreachable
4845 /// servers meant a full round of connect timeouts before every first token —
4846 /// the "MCP is always reloading and slowing things down" report. The second
4847 /// pass must produce the same diagnosis without dialing anything.
4848 #[tokio::test]
4849 async fn a_failed_server_waits_out_a_cooldown_instead_of_redialing_every_turn() {
4850 let dir = tempfile::tempdir().unwrap();
4851 let path = dir.path().join("mcp.json");
4852 fs::write(
4853 &path,
4854 r#"{
4855 "mcpServers": {
4856 "broken": {
4857 "command": "codewhale-tui-test-this-binary-does-not-exist-9f8e7d6c5b4a",
4858 "args": []
4859 }
4860 }
4861 }"#,
4862 )
4863 .unwrap();
4864
4865 let mut pool = McpPool::from_config_path(&path).unwrap();
4866 let first = pool.connect_all().await;
4867 assert_eq!(first.len(), 1, "first pass should dial and fail once");
4868
4869 // Second pass: still reported as failing, but nothing is queued to dial.
4870 let (pending, errors) = pool.collect_pending_connects(None);
4871 assert!(
4872 pending.is_empty(),
4873 "a server inside its cooldown must not be re-dialed: {:?}",
4874 pending.iter().map(|(name, _)| name).collect::<Vec<_>>()
4875 );
4876 assert_eq!(errors.len(), 1, "the failure must still be reported");
4877 assert_eq!(errors[0].0, "broken");
4878 assert!(
4879 format!("{:#}", errors[0].1)
4880 .to_lowercase()
4881 .contains("spawn"),
4882 "the replayed diagnosis must be the real one: {:#}",
4883 errors[0].1
4884 );
4885
4886 // Asking for that server by name is explicit intent and lifts the wait.
4887 assert!(pool.retry_connection("broken").await.is_err());
4888 let (pending, _) = pool.collect_pending_connects(None);
4889 assert!(
4890 pending.is_empty(),
4891 "the failed retry restarts the ladder rather than clearing it"
4892 );
4893 }
4894
4895 /// Lazy boot (#6033): the scoped collect starts only the eager set —
4896 /// `required` servers plus ones an explicit tool selection covers — and the
4897 /// pool tracks exactly those names as in-flight, so "connecting" never has
4898 /// to be inferred from "enabled but unconnected".
4899 #[test]
4900 fn lazy_boot_scopes_pending_connects_and_tracks_in_flight() {
4901 let mut required_cfg = test_server_config();
4902 required_cfg.required = true;
4903 let mut pool = McpPool::new(McpConfig {
4904 timeouts: McpTimeouts::default(),
4905 servers: HashMap::from([
4906 ("needed".to_string(), required_cfg),
4907 ("selected".to_string(), test_server_config()),
4908 ("lazy".to_string(), test_server_config()),
4909 ]),
4910 });
4911 let requested = vec!["mcp_selected_read".to_string()];
4912
4913 let eager = pool.eager_boot_server_names(&requested);
4914 assert_eq!(
4915 eager,
4916 HashSet::from(["needed".to_string(), "selected".to_string()]),
4917 "the eager set is required servers plus selection-covered ones"
4918 );
4919
4920 let (pending, errors) = pool.collect_pending_connects(Some(&eager));
4921 assert!(errors.is_empty());
4922 let pending_names: HashSet<String> = pending.iter().map(|(name, _)| name.clone()).collect();
4923 assert_eq!(pending_names, eager);
4924 assert_eq!(
4925 pool.connecting_servers()
4926 .into_iter()
4927 .collect::<HashSet<_>>(),
4928 pending_names,
4929 "in-flight marks must name exactly the spawned connects"
4930 );
4931
4932 // A lazy server is neither spawned nor reported connecting.
4933 let (pending, errors) = pool.collect_pending_connects(Some(&eager));
4934 assert!(
4935 pending.is_empty() && errors.is_empty(),
4936 "an in-flight name is not re-queued by a second scoped pass"
4937 );
4938
4939 // Explicit selection is intent: the lazy server starts on demand and is
4940 // marked in-flight while it does.
4941 let (pending, errors) = pool.take_pending_connects_for(&["lazy".to_string()]);
4942 assert!(errors.is_empty());
4943 assert_eq!(pending.len(), 1);
4944 assert!(pool.connecting_servers().contains(&"lazy".to_string()));
4945
4946 // Aborting clears the marks without touching connection state.
4947 pool.cancel_connecting(&HashSet::from([
4948 "needed".to_string(),
4949 "selected".to_string(),
4950 "lazy".to_string(),
4951 ]));
4952 assert!(pool.connecting_servers().is_empty());
4953 }
4954
4955 /// Selection coverage shared by lazy boot and the per-turn wait: exact
4956 /// `mcp_<server>_<tool>` names and `mcp_<prefix>*` globs both count.
4957 #[test]
4958 fn tool_selection_covers_exact_names_and_globs() {
4959 let selected = vec![
4960 "mcp_fs_read".to_string(),
4961 "mcp_git_*".to_string(),
4962 "shell".to_string(),
4963 ];
4964 assert!(tool_selection_covers_server(&selected, "fs"));
4965 assert!(tool_selection_covers_server(&selected, "git_status"));
4966 assert!(!tool_selection_covers_server(&selected, "slack"));
4967 // A prefix glob reaches every server whose `mcp_<server>_` namespace
4968 // starts with it: `mcp_gi*` covers `git` and `gitea` alike.
4969 let glob = vec!["mcp_gi*".to_string()];
4970 assert!(tool_selection_covers_server(&glob, "git"));
4971 assert!(tool_selection_covers_server(&glob, "gitea"));
4972 assert!(!tool_selection_covers_server(&glob, "fs"));
4973 }
4974
4975 #[test]
4976 fn connect_backoff_doubles_then_holds_at_the_ceiling() {
4977 use std::time::Duration;
4978 assert_eq!(connect_backoff_delay(1), Duration::from_secs(30));
4979 assert_eq!(connect_backoff_delay(2), Duration::from_secs(60));
4980 assert_eq!(connect_backoff_delay(3), Duration::from_secs(120));
4981 assert_eq!(connect_backoff_delay(6), Duration::from_secs(600));
4982 assert_eq!(
4983 connect_backoff_delay(50),
4984 Duration::from_secs(600),
4985 "the ceiling holds; a long-dead server is retried every ten minutes"
4986 );
4987 }
4988
4989 #[test]
4990 fn parse_sse_message_data_extracts_message_events() {
4991 let body = "event: message\r\ndata: {\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\r\n\r\n";
4992 let messages = parse_sse_message_data(body);
4993 assert_eq!(messages.len(), 1);
4994 let value: serde_json::Value = serde_json::from_slice(&messages[0]).unwrap();
4995 assert_eq!(value["id"], 1);
4996 assert!(value.get("result").is_some());
4997 }
4998
4999 #[test]
5000 fn response_id_matches_string_and_numeric_echoes() {
5001 assert!(response_id_matches(Some(&serde_json::json!("1")), "1"));
5002 assert!(response_id_matches(Some(&serde_json::json!(1)), "1"));
5003 assert!(!response_id_matches(Some(&serde_json::json!("2")), "1"));
5004 }
5005
5006 #[test]
5007 fn legacy_sse_transport_requires_explicit_config() {
5008 let mut server = test_server_config();
5009 server.url = Some("https://example.com/mcp/abc/sse".to_string());
5010
5011 assert!(
5012 !is_legacy_sse_transport(&server),
5013 "/sse paths must not force legacy SSE without an explicit transport override"
5014 );
5015
5016 server.transport = Some("sse".to_string());
5017 assert!(is_legacy_sse_transport(&server));
5018
5019 server.transport = Some("SSE".to_string());
5020 assert!(is_legacy_sse_transport(&server));
5021
5022 server.transport = Some("http".to_string());
5023 assert!(!is_legacy_sse_transport(&server));
5024 }
5025
5026 #[test]
5027 fn find_sse_event_separator_accepts_lf_and_crlf() {
5028 assert_eq!(
5029 find_sse_event_separator("event: endpoint\n\n"),
5030 Some((15, 2))
5031 );
5032 assert_eq!(
5033 find_sse_event_separator("event: endpoint\r\n\r\n"),
5034 Some((15, 4))
5035 );
5036 }
5037
5038 #[test]
5039 fn find_sse_event_separator_bytes_matches_str_and_survives_multibyte() {
5040 // Same offsets as the str version.
5041 assert_eq!(
5042 find_sse_event_separator_bytes(b"event: endpoint\n\n"),
5043 Some((15, 2))
5044 );
5045 assert_eq!(
5046 find_sse_event_separator_bytes(b"event: endpoint\r\n\r\n"),
5047 Some((15, 4))
5048 );
5049 // A frame whose data holds a multi-byte char, accumulated byte-wise and
5050 // split mid-char across two reads, decodes intact (no U+FFFD).
5051 let frame = "data: 你好\n\n";
5052 let bytes = frame.as_bytes();
5053 let split = bytes.len() - 3; // inside "好" / before the separator
5054 let mut buffer: Vec<u8> = Vec::new();
5055 buffer.extend_from_slice(&bytes[..split]);
5056 assert_eq!(find_sse_event_separator_bytes(&buffer), None);
5057 buffer.extend_from_slice(&bytes[split..]);
5058 let (pos, sep) = find_sse_event_separator_bytes(&buffer).expect("separator");
5059 let block = String::from_utf8_lossy(&buffer[..pos]).into_owned();
5060 assert_eq!(block, "data: 你好");
5061 assert!(!block.contains('\u{FFFD}'), "multibyte corrupted");
5062 assert_eq!(sep, 2);
5063 }
5064
5065 #[tokio::test]
5066 async fn mcp_connection_supports_streamable_http_event_stream_responses() {
5067 use tokio::io::{AsyncReadExt, AsyncWriteExt};
5068 use tokio::net::{TcpListener, TcpStream};
5069
5070 async fn read_http_request(socket: &mut TcpStream) -> String {
5071 let mut request = Vec::new();
5072 let mut buf = [0; 1024];
5073 let header_end = loop {
5074 let n = socket.read(&mut buf).await.unwrap();
5075 assert!(n > 0, "client closed before headers completed");
5076 request.extend_from_slice(&buf[..n]);
5077 if let Some(pos) = request.windows(4).position(|window| window == b"\r\n\r\n") {
5078 break pos + 4;
5079 }
5080 };
5081
5082 let headers = String::from_utf8_lossy(&request[..header_end]);
5083 let content_length = headers
5084 .lines()
5085 .find_map(|line| {
5086 let (name, value) = line.split_once(':')?;
5087 name.eq_ignore_ascii_case("content-length")
5088 .then(|| value.trim().parse::<usize>().ok())
5089 .flatten()
5090 })
5091 .unwrap_or(0);
5092 let total_len = header_end + content_length;
5093 while request.len() < total_len {
5094 let n = socket.read(&mut buf).await.unwrap();
5095 assert!(n > 0, "client closed before body completed");
5096 request.extend_from_slice(&buf[..n]);
5097 }
5098
5099 String::from_utf8(request).unwrap()
5100 }
5101
5102 async fn write_json_sse(socket: &mut TcpStream, response: serde_json::Value) {
5103 let body = format!("event: message\ndata: {response}\n\n");
5104 let response = format!(
5105 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\n\r\n{}",
5106 body.len(),
5107 body
5108 );
5109 socket.write_all(response.as_bytes()).await.unwrap();
5110 }
5111
5112 let _lock = lock_mcp_loopback_tests().await;
5113 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5114 let addr = listener.local_addr().unwrap();
5115 let server = tokio::spawn(async move {
5116 loop {
5117 let Ok((mut socket, _)) = listener.accept().await else {
5118 break;
5119 };
5120 tokio::spawn(async move {
5121 let request = read_http_request(&mut socket).await;
5122 assert!(request.starts_with("POST /mcp "));
5123 assert!(
5124 request.contains("Accept: application/json, text/event-stream")
5125 || request.contains("accept: application/json, text/event-stream")
5126 );
5127 let body = request.split("\r\n\r\n").nth(1).unwrap_or("");
5128 let value: serde_json::Value = serde_json::from_str(body).unwrap();
5129 let method = value["method"].as_str().unwrap();
5130
5131 if method == "notifications/initialized" {
5132 socket
5133 .write_all(b"HTTP/1.1 202 Accepted\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
5134 .await
5135 .unwrap();
5136 return;
5137 }
5138
5139 let id = value["id"].clone();
5140 let result = match method {
5141 "initialize" => serde_json::json!({
5142 "protocolVersion": "2024-11-05",
5143 "serverInfo": {"name": "mock-streamable", "version": "1.0.0"},
5144 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
5145 }),
5146 "tools/list" => serde_json::json!({
5147 "tools": [{
5148 "name": "read_wiki_structure",
5149 "description": "Read wiki structure",
5150 "inputSchema": {"type": "object"}
5151 }]
5152 }),
5153 "resources/list" => serde_json::json!({"resources": []}),
5154 "resources/templates/list" => {
5155 serde_json::json!({"resourceTemplates": []})
5156 }
5157 "prompts/list" => serde_json::json!({"prompts": []}),
5158 other => panic!("unexpected method: {other}"),
5159 };
5160 write_json_sse(
5161 &mut socket,
5162 serde_json::json!({
5163 "jsonrpc": "2.0",
5164 "id": id,
5165 "result": result
5166 }),
5167 )
5168 .await;
5169 });
5170 }
5171 });
5172
5173 let config = McpServerConfig {
5174 command: None,
5175 args: vec![],
5176 env: HashMap::new(),
5177 cwd: None,
5178 url: Some(format!("http://{addr}/mcp")),
5179 transport: None,
5180 connect_timeout: Some(5),
5181 execute_timeout: None,
5182 read_timeout: None,
5183 disabled: false,
5184 enabled: true,
5185 required: false,
5186 enabled_tools: Vec::new(),
5187 disabled_tools: Vec::new(),
5188 headers: HashMap::new(),
5189 env_headers: HashMap::new(),
5190 bearer_token_env_var: None,
5191 scopes: Vec::new(),
5192 oauth: None,
5193 oauth_resource: None,
5194 reviewed_plugin: None,
5195 runtime_added: false,
5196 allow_private_network: false,
5197 };
5198
5199 let conn = McpConnection::connect_with_policy(
5200 "deepwiki".to_string(),
5201 config,
5202 &McpTimeouts::default(),
5203 None,
5204 )
5205 .await
5206 .unwrap();
5207
5208 assert_eq!(conn.state(), ConnectionState::Ready);
5209 assert_eq!(conn.tools().len(), 1);
5210 assert_eq!(conn.tools()[0].name, "read_wiki_structure");
5211
5212 server.abort();
5213 }
5214
5215 #[test]
5216 fn mask_url_secrets_strips_userinfo() {
5217 let masked = mask_url_secrets("https://user:s3cret@host.example/api?foo=bar");
5218 assert!(masked.contains("***"), "expected masked userinfo: {masked}");
5219 assert!(!masked.contains("s3cret"), "secret leaked: {masked}");
5220 assert!(masked.contains("host.example"), "host preserved: {masked}");
5221 }
5222
5223 #[test]
5224 fn mask_url_secrets_passes_through_clean_url() {
5225 assert_eq!(
5226 mask_url_secrets("https://api.example.com/mcp"),
5227 "https://api.example.com/mcp"
5228 );
5229 }
5230
5231 #[test]
5232 fn redact_body_preview_masks_bearer_token() {
5233 let redacted = redact_body_preview(
5234 "Authorization: Bearer abc.def.ghi end; authorization: bearer second-token end",
5235 );
5236 assert_eq!(
5237 redacted.matches("Bearer ***").count() + redacted.matches("bearer ***").count(),
5238 2,
5239 "redacted: {redacted}"
5240 );
5241 assert!(
5242 !redacted.contains("abc.def.ghi") && !redacted.contains("second-token"),
5243 "leaked: {redacted}"
5244 );
5245 }
5246
5247 #[test]
5248 fn redact_proxy_userinfo_strips_password() {
5249 // Corporate-style proxy URL with embedded creds — the
5250 // password must never reach the on-disk log file. URL strings
5251 // are assembled from placeholder constants via `format!` so the
5252 // literal source never contains a scheme-prefixed username +
5253 // password pair (colon-separated, `@`-terminated) that
5254 // GitGuardian's "Basic Auth String" detector would flag as a
5255 // committed credential.
5256 let (placeholder_user, placeholder_pass) = ("PLACEHOLDER_USER", "PLACEHOLDER_PASS");
5257 let with_creds = format!("http://{placeholder_user}:{placeholder_pass}@proxy.example/");
5258 let redacted = redact_proxy_userinfo(&with_creds);
5259 assert_eq!(redacted, "http://***@proxy.example/");
5260 assert!(!redacted.contains(placeholder_pass));
5261 assert!(!redacted.contains(placeholder_user));
5262
5263 // User only (no password) — still redacted.
5264 let with_user_only = format!("https://{placeholder_user}@proxy.example:8080");
5265 let redacted = redact_proxy_userinfo(&with_user_only);
5266 assert_eq!(redacted, "https://***@proxy.example:8080");
5267
5268 // No userinfo segment — pass through.
5269 let redacted = redact_proxy_userinfo("http://proxy.example:3128/");
5270 assert_eq!(redacted, "http://proxy.example:3128/");
5271
5272 // `@` appears only in the path, not as userinfo separator —
5273 // must not be mistaken for credentials.
5274 let redacted = redact_proxy_userinfo("http://proxy.example/path@thing");
5275 assert_eq!(redacted, "http://proxy.example/path@thing");
5276
5277 // Garbage input (no `://`) returned unchanged — the
5278 // surrounding warning log is the only caller and is already
5279 // handling the malformed-URL case.
5280 assert_eq!(redact_proxy_userinfo("not-a-url"), "not-a-url");
5281 }
5282
5283 #[test]
5284 fn redact_body_preview_masks_api_key_param() {
5285 let redacted = redact_body_preview("error api_key=sk-12345&other=val then TOKEN=second-secret");
5286 assert!(redacted.contains("api_key=***"), "redacted: {redacted}");
5287 assert!(redacted.contains("TOKEN=***"), "redacted: {redacted}");
5288 assert!(
5289 !redacted.contains("sk-12345") && !redacted.contains("second-secret"),
5290 "leaked: {redacted}"
5291 );
5292 assert!(
5293 redacted.contains("other=val"),
5294 "non-secret preserved: {redacted}"
5295 );
5296 }
5297
5298 #[test]
5299 fn reviewed_plugin_server_errors_suppress_arbitrary_details() {
5300 let auth = McpHttpAuth {
5301 suppress_server_error_details: true,
5302 ..Default::default()
5303 };
5304 assert_eq!(
5305 auth.server_error_preview("arbitrary credential value"),
5306 "<server details suppressed for reviewed plugin>"
5307 );
5308
5309 let response = serde_json::json!({
5310 "error": { "message": "arbitrary credential value" }
5311 });
5312 let error = response_result(&response, "tools/call", true)
5313 .expect_err("reviewed plugin JSON-RPC error must be generic")
5314 .to_string();
5315 assert!(!error.contains("arbitrary credential value"));
5316 assert!(error.contains("details suppressed"));
5317 }
5318
5319 #[test]
5320 fn invalid_json_preview_collapses_lines_and_redacts_secrets() {
5321 let preview = invalid_json_preview(
5322 b"Authorization: Bearer PLACEHOLDER_TOKEN\nAllow connection? api_key=PLACEHOLDER_KEY",
5323 );
5324
5325 assert!(
5326 preview.contains("Authorization: Bearer *** Allow connection? api_key=***"),
5327 "preview: {preview}"
5328 );
5329 assert!(
5330 !preview.contains('\n'),
5331 "preview should be single-line: {preview}"
5332 );
5333 assert!(
5334 !preview.contains("PLACEHOLDER_TOKEN") && !preview.contains("PLACEHOLDER_KEY"),
5335 "secret leaked: {preview}"
5336 );
5337 }
5338
5339 /// #420: `StdioTransport::shutdown` reaps the child process by sending
5340 /// SIGTERM and giving it a brief grace period before drop fires SIGKILL.
5341 /// The test spawns `cat` (which exits immediately on stdin EOF / SIGTERM)
5342 /// and verifies the transport tears down cleanly. Unix-only because
5343 /// SIGTERM doesn't exist on Windows; on Windows the test would just
5344 /// duplicate the kill_on_drop path.
5345 #[cfg(unix)]
5346 #[tokio::test]
5347 async fn stdio_transport_shutdown_terminates_child() {
5348 let mut config = test_server_config();
5349 config.command = Some("cat".to_string());
5350 let mut transport = StdioTransport::spawn(
5351 "shutdown-test",
5352 "cat",
5353 &config,
5354 tokio_util::sync::CancellationToken::new(),
5355 )
5356 .expect("spawn cat through the production broker");
5357 let pid = transport
5358 .session
5359 .child_for_tests()
5360 .lock()
5361 .await
5362 .id()
5363 .expect("child pid");
5364
5365 // shutdown() should send SIGTERM and complete within the grace window.
5366 let start = std::time::Instant::now();
5367 transport.shutdown().await;
5368 let elapsed = start.elapsed();
5369 assert!(
5370 elapsed < STDIO_SHUTDOWN_GRACE + Duration::from_millis(500),
5371 "shutdown blocked beyond grace window: {elapsed:?}"
5372 );
5373
5374 // The child should be reaped — kill(pid, 0) returning ESRCH means
5375 // the pid is gone. If it's still alive, kill(0) returns 0, which
5376 // means our shutdown didn't terminate it.
5377 // SAFETY: pid was just collected from a tokio Child we spawned.
5378 // libc::kill with signal 0 only checks pid existence and is
5379 // async-signal-safe.
5380 let still_alive = unsafe { libc::kill(pid as i32, 0) } == 0;
5381 assert!(
5382 !still_alive,
5383 "child {pid} survived StdioTransport::shutdown — SIGTERM not delivered"
5384 );
5385 }
5386
5387 #[cfg(unix)]
5388 #[tokio::test]
5389 async fn stdio_transport_drop_allows_child_cleanup() {
5390 let directory = tempfile::tempdir().expect("temporary cleanup receipt");
5391 let receipt = directory.path().join("cleaned");
5392 let config: McpServerConfig = serde_json::from_value(serde_json::json!({
5393 "args": [
5394 "-c",
5395 "trap 'sleep 0.1; printf cleaned > \"$1.tmp\"; mv \"$1.tmp\" \"$1\"; exit 0' TERM; printf 'ready\\n'; while :; do sleep 0.05; done",
5396 "cleanup-test",
5397 receipt.display().to_string(),
5398 ],
5399 }))
5400 .unwrap();
5401 let mut transport = StdioTransport::spawn(
5402 "drop-cleanup-test",
5403 "/bin/sh",
5404 &config,
5405 tokio_util::sync::CancellationToken::new(),
5406 )
5407 .expect("spawn cleanup fixture");
5408 assert_eq!(transport.recv().await.unwrap(), b"ready");
5409 // Retaining a strong Child reference would mask immediate kill_on_drop.
5410 drop(transport);
5411 tokio::time::timeout(STDIO_SHUTDOWN_GRACE + Duration::from_secs(1), async {
5412 while !receipt.exists() {
5413 tokio::time::sleep(Duration::from_millis(20)).await;
5414 }
5415 })
5416 .await
5417 .expect("dropped transport lets SIGTERM cleanup finish");
5418 assert_eq!(std::fs::read_to_string(receipt).unwrap(), "cleaned");
5419 }
5420
5421 #[cfg(unix)]
5422 #[tokio::test]
5423 async fn stdio_transport_drop_kills_child_that_ignores_cleanup() {
5424 let config: McpServerConfig = serde_json::from_value(serde_json::json!({
5425 "args": [
5426 "-c",
5427 "trap '' TERM; printf 'ready\\n'; while :; do sleep 0.05; done",
5428 ],
5429 }))
5430 .unwrap();
5431 let mut transport = StdioTransport::spawn(
5432 "drop-hung-test",
5433 "/bin/sh",
5434 &config,
5435 tokio_util::sync::CancellationToken::new(),
5436 )
5437 .expect("spawn unresponsive fixture");
5438 assert_eq!(transport.recv().await.unwrap(), b"ready");
5439 let pid = transport
5440 .session
5441 .child_for_tests()
5442 .lock()
5443 .await
5444 .id()
5445 .expect("live child");
5446 drop(transport);
5447 tokio::time::timeout(STDIO_SHUTDOWN_GRACE + Duration::from_secs(1), async {
5448 // Signal zero only observes the process; the owned Child sends kills.
5449 while unsafe { libc::kill(pid as i32, 0) } == 0 {
5450 tokio::time::sleep(Duration::from_millis(20)).await;
5451 }
5452 })
5453 .await
5454 .expect("dropped transport force-kills and reaps hung child");
5455 }
5456
5457 /// Mid-run MCP server crash: the v0.8.x spawn path used `Stdio::null` for
5458 /// stderr, so a server that died with a useful stderr message left the
5459 /// caller with only "Stdio transport closed". Now stderr is piped into a
5460 /// bounded ring buffer and surfaced when the read side fails.
5461 #[cfg(unix)]
5462 #[tokio::test]
5463 async fn stdio_transport_recv_error_includes_stderr_tail() {
5464 let mut config = test_server_config();
5465 config.command = Some("sh".to_string());
5466 config.args = vec![
5467 "-c".to_string(),
5468 "echo 'mcp-server: failed to load plugin' 1>&2; exit 1".to_string(),
5469 ];
5470 let mut transport = StdioTransport::spawn(
5471 "stderr-test",
5472 "sh",
5473 &config,
5474 tokio_util::sync::CancellationToken::new(),
5475 )
5476 .expect("spawn sh through the production broker");
5477
5478 // Give the subprocess time to write its stderr line and exit.
5479 tokio::time::sleep(Duration::from_millis(300)).await;
5480
5481 let err = transport
5482 .recv()
5483 .await
5484 .expect_err("expected transport closed error");
5485 let err_str = format!("{err}");
5486 assert!(
5487 err_str.contains("Stdio transport closed"),
5488 "missing closed marker in: {err_str}"
5489 );
5490 assert!(
5491 err_str.contains("mcp-server: failed to load plugin"),
5492 "stderr context missing from error: {err_str}"
5493 );
5494 }
5495
5496 #[tokio::test]
5497 async fn sse_connect_waits_for_endpoint_before_first_send() {
5498 use std::sync::{
5499 Arc,
5500 atomic::{AtomicBool, Ordering as AtomicOrdering},
5501 };
5502 use tokio::io::{AsyncReadExt, AsyncWriteExt};
5503 use tokio::net::TcpListener;
5504
5505 let _lock = lock_mcp_loopback_tests().await;
5506 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5507 let addr = listener.local_addr().unwrap();
5508 let post_seen = Arc::new(AtomicBool::new(false));
5509 let server_post_seen = Arc::clone(&post_seen);
5510 let cancel_token = tokio_util::sync::CancellationToken::new();
5511 let server_cancel = cancel_token.clone();
5512
5513 let server = tokio::spawn(async move {
5514 loop {
5515 let Ok((mut socket, _)) = listener.accept().await else {
5516 break;
5517 };
5518 let post_seen = Arc::clone(&server_post_seen);
5519 let server_cancel = server_cancel.clone();
5520 tokio::spawn(async move {
5521 let mut request = Vec::new();
5522 let mut buf = [0; 1024];
5523 loop {
5524 let n = socket.read(&mut buf).await.unwrap();
5525 if n == 0 {
5526 return;
5527 }
5528 request.extend_from_slice(&buf[..n]);
5529 if request.windows(4).any(|window| window == b"\r\n\r\n") {
5530 break;
5531 }
5532 }
5533 let request = String::from_utf8_lossy(&request);
5534 if request.starts_with("GET /sse ") {
5535 socket
5536 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
5537 .await
5538 .unwrap();
5539 tokio::time::sleep(Duration::from_millis(150)).await;
5540 socket
5541 .write_all(b"event: endpoint\ndata: /messages\n\n")
5542 .await
5543 .unwrap();
5544 server_cancel.cancelled().await;
5545 } else if request.starts_with("POST /messages ") {
5546 post_seen.store(true, AtomicOrdering::SeqCst);
5547 socket
5548 .write_all(
5549 b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
5550 )
5551 .await
5552 .unwrap();
5553 }
5554 });
5555 }
5556 });
5557
5558 let url = format!("http://{addr}/sse");
5559 let client = test_mcp_http_client(&url);
5560 let mut transport =
5561 SseTransport::connect(client, url, cancel_token.clone(), Duration::from_secs(2))
5562 .await
5563 .unwrap();
5564
5565 transport
5566 .send(json_frame(serde_json::json!({
5567 "jsonrpc": "2.0",
5568 "id": 1,
5569 "method": "initialize"
5570 })))
5571 .await
5572 .unwrap();
5573
5574 assert!(
5575 post_seen.load(AtomicOrdering::SeqCst),
5576 "first SSE send should POST to the discovered endpoint"
5577 );
5578
5579 cancel_token.cancel();
5580 server.abort();
5581 }
5582
5583 #[tokio::test]
5584 async fn sse_connect_accepts_crlf_endpoint_events() {
5585 use std::sync::{
5586 Arc,
5587 atomic::{AtomicBool, Ordering as AtomicOrdering},
5588 };
5589 use tokio::io::{AsyncReadExt, AsyncWriteExt};
5590 use tokio::net::TcpListener;
5591
5592 let _lock = lock_mcp_loopback_tests().await;
5593 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5594 let addr = listener.local_addr().unwrap();
5595 let post_seen = Arc::new(AtomicBool::new(false));
5596 let server_post_seen = Arc::clone(&post_seen);
5597 let cancel_token = tokio_util::sync::CancellationToken::new();
5598 let server_cancel = cancel_token.clone();
5599
5600 let server = tokio::spawn(async move {
5601 loop {
5602 let Ok((mut socket, _)) = listener.accept().await else {
5603 break;
5604 };
5605 let post_seen = Arc::clone(&server_post_seen);
5606 let server_cancel = server_cancel.clone();
5607 tokio::spawn(async move {
5608 let mut request = Vec::new();
5609 let mut buf = [0; 1024];
5610 loop {
5611 let n = socket.read(&mut buf).await.unwrap();
5612 if n == 0 {
5613 return;
5614 }
5615 request.extend_from_slice(&buf[..n]);
5616 if request.windows(4).any(|window| window == b"\r\n\r\n") {
5617 break;
5618 }
5619 }
5620 let request = String::from_utf8_lossy(&request);
5621 if request.starts_with("GET /sse ") {
5622 socket
5623 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
5624 .await
5625 .unwrap();
5626 socket
5627 .write_all(b"event: endpoint\r\ndata: /messages\r\n\r\n")
5628 .await
5629 .unwrap();
5630 server_cancel.cancelled().await;
5631 } else if request.starts_with("POST /messages ") {
5632 post_seen.store(true, AtomicOrdering::SeqCst);
5633 socket
5634 .write_all(
5635 b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
5636 )
5637 .await
5638 .unwrap();
5639 }
5640 });
5641 }
5642 });
5643
5644 let url = format!("http://{addr}/sse");
5645 let client = test_mcp_http_client(&url);
5646 let mut transport =
5647 SseTransport::connect(client, url, cancel_token.clone(), Duration::from_secs(2))
5648 .await
5649 .unwrap();
5650
5651 transport
5652 .send(json_frame(serde_json::json!({
5653 "jsonrpc": "2.0",
5654 "id": 1,
5655 "method": "initialize"
5656 })))
5657 .await
5658 .unwrap();
5659
5660 assert!(
5661 post_seen.load(AtomicOrdering::SeqCst),
5662 "first SSE send should POST to the CRLF-discovered endpoint"
5663 );
5664
5665 cancel_token.cancel();
5666 server.abort();
5667 }
5668
5669 #[tokio::test]
5670 async fn sse_transport_applies_custom_headers_to_get_and_post() {
5671 use std::sync::{
5672 Arc,
5673 atomic::{AtomicBool, Ordering as AtomicOrdering},
5674 };
5675 use tokio::io::{AsyncReadExt, AsyncWriteExt};
5676 use tokio::net::TcpListener;
5677
5678 let _lock = lock_mcp_loopback_tests().await;
5679 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5680 let addr = listener.local_addr().unwrap();
5681 let get_header_seen = Arc::new(AtomicBool::new(false));
5682 let post_header_seen = Arc::new(AtomicBool::new(false));
5683 let server_get_header_seen = Arc::clone(&get_header_seen);
5684 let server_post_header_seen = Arc::clone(&post_header_seen);
5685 let cancel_token = tokio_util::sync::CancellationToken::new();
5686 let server_cancel = cancel_token.clone();
5687
5688 let server = tokio::spawn(async move {
5689 loop {
5690 let Ok((mut socket, _)) = listener.accept().await else {
5691 break;
5692 };
5693 let get_header_seen = Arc::clone(&server_get_header_seen);
5694 let post_header_seen = Arc::clone(&server_post_header_seen);
5695 let server_cancel = server_cancel.clone();
5696 tokio::spawn(async move {
5697 let mut request = Vec::new();
5698 let mut buf = [0; 1024];
5699 loop {
5700 let n = socket.read(&mut buf).await.unwrap();
5701 if n == 0 {
5702 return;
5703 }
5704 request.extend_from_slice(&buf[..n]);
5705 if request.windows(4).any(|window| window == b"\r\n\r\n") {
5706 break;
5707 }
5708 }
5709 let request = String::from_utf8_lossy(&request);
5710 let request_lower = request.to_lowercase();
5711 if request.starts_with("GET /sse ") {
5712 if request_lower.contains("x-custom-auth: my-test-token") {
5713 get_header_seen.store(true, AtomicOrdering::SeqCst);
5714 }
5715 socket
5716 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
5717 .await
5718 .unwrap();
5719 socket
5720 .write_all(b"event: endpoint\ndata: /messages\n\n")
5721 .await
5722 .unwrap();
5723 server_cancel.cancelled().await;
5724 } else if request.starts_with("POST /messages ") {
5725 if request_lower.contains("x-custom-auth: my-test-token") {
5726 post_header_seen.store(true, AtomicOrdering::SeqCst);
5727 }
5728 socket
5729 .write_all(
5730 b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
5731 )
5732 .await
5733 .unwrap();
5734 }
5735 });
5736 }
5737 });
5738
5739 let url = format!("http://{addr}/sse");
5740 let client = test_mcp_http_client(&url);
5741 let mut headers = HashMap::new();
5742 headers.insert("X-Custom-Auth".to_string(), "my-test-token".to_string());
5743 let mut transport = SseTransport::connect(
5744 client.with_mcp_auth(McpHttpAuth {
5745 headers,
5746 ..Default::default()
5747 }),
5748 url,
5749 cancel_token.clone(),
5750 Duration::from_secs(2),
5751 )
5752 .await
5753 .unwrap();
5754
5755 transport
5756 .send(json_frame(serde_json::json!({
5757 "jsonrpc": "2.0",
5758 "id": 1,
5759 "method": "initialize"
5760 })))
5761 .await
5762 .unwrap();
5763
5764 assert!(
5765 get_header_seen.load(AtomicOrdering::SeqCst),
5766 "legacy SSE GET must include user-configured custom headers"
5767 );
5768 assert!(
5769 post_header_seen.load(AtomicOrdering::SeqCst),
5770 "legacy SSE POST must include user-configured custom headers"
5771 );
5772
5773 cancel_token.cancel();
5774 server.abort();
5775 }
5776
5777 #[tokio::test]
5778 async fn sse_post_error_includes_response_body_excerpt() {
5779 use tokio::io::{AsyncReadExt, AsyncWriteExt};
5780 use tokio::net::TcpListener;
5781
5782 let _lock = lock_mcp_loopback_tests().await;
5783 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5784 let addr = listener.local_addr().unwrap();
5785 let cancel_token = tokio_util::sync::CancellationToken::new();
5786 let server_cancel = cancel_token.clone();
5787
5788 let server = tokio::spawn(async move {
5789 loop {
5790 let Ok((mut socket, _)) = listener.accept().await else {
5791 break;
5792 };
5793 let server_cancel = server_cancel.clone();
5794 tokio::spawn(async move {
5795 let mut request = Vec::new();
5796 let mut buf = [0; 1024];
5797 loop {
5798 let n = socket.read(&mut buf).await.unwrap();
5799 if n == 0 {
5800 return;
5801 }
5802 request.extend_from_slice(&buf[..n]);
5803 if request.windows(4).any(|window| window == b"\r\n\r\n") {
5804 break;
5805 }
5806 }
5807 let request = String::from_utf8_lossy(&request);
5808 if request.starts_with("GET /sse ") {
5809 socket
5810 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
5811 .await
5812 .unwrap();
5813 socket
5814 .write_all(b"event: endpoint\ndata: /messages\n\n")
5815 .await
5816 .unwrap();
5817 server_cancel.cancelled().await;
5818 } else if request.starts_with("POST /messages ") {
5819 socket
5820 .write_all(
5821 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\"}",
5822 )
5823 .await
5824 .unwrap();
5825 }
5826 });
5827 }
5828 });
5829
5830 let url = format!("http://{addr}/sse");
5831 let client = test_mcp_http_client(&url);
5832 let mut transport =
5833 SseTransport::connect(client, url, cancel_token.clone(), Duration::from_secs(2))
5834 .await
5835 .unwrap();
5836
5837 let err = transport
5838 .send(json_frame(serde_json::json!({
5839 "jsonrpc": "2.0",
5840 "id": 1,
5841 "method": "initialize"
5842 })))
5843 .await
5844 .expect_err("POST rejection should be returned");
5845 let err = format!("{err:#}");
5846 assert!(
5847 err.contains("400 Bad Request") && err.contains("missing query"),
5848 "SSE POST error should include status and body, got: {err}"
5849 );
5850
5851 cancel_token.cancel();
5852 server.abort();
5853 }
5854
5855 #[tokio::test]
5856 async fn streamable_http_caps_chunked_bodies_without_content_length() {
5857 use tokio::io::{AsyncReadExt, AsyncWriteExt};
5858 use tokio::net::TcpListener;
5859
5860 let _lock = lock_mcp_loopback_tests().await;
5861 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5862 let addr = listener.local_addr().unwrap();
5863
5864 // Serve chunked responses (no Content-Length) of the requested size:
5865 // GET /over streams past the cap, GET /under stays below it.
5866 let server = tokio::spawn(async move {
5867 loop {
5868 let Ok((mut socket, _)) = listener.accept().await else {
5869 break;
5870 };
5871 tokio::spawn(async move {
5872 let mut request = Vec::new();
5873 let mut buf = [0; 1024];
5874 loop {
5875 let n = socket.read(&mut buf).await.unwrap();
5876 if n == 0 {
5877 return;
5878 }
5879 request.extend_from_slice(&buf[..n]);
5880 if request.windows(4).any(|window| window == b"\r\n\r\n") {
5881 break;
5882 }
5883 }
5884 let request = String::from_utf8_lossy(&request);
5885 let total: usize = if request.starts_with("GET /over ") {
5886 256
5887 } else {
5888 16
5889 };
5890 socket
5891 .write_all(
5892 b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nTransfer-Encoding: chunked\r\n\r\n",
5893 )
5894 .await
5895 .unwrap();
5896 let chunk = [b'x'; 32];
5897 let mut sent = 0;
5898 while sent < total {
5899 let n = chunk.len().min(total - sent);
5900 let frame = format!("{n:x}\r\n");
5901 socket.write_all(frame.as_bytes()).await.unwrap();
5902 socket.write_all(&chunk[..n]).await.unwrap();
5903 socket.write_all(b"\r\n").await.unwrap();
5904 sent += n;
5905 }
5906 socket.write_all(b"0\r\n\r\n").await.unwrap();
5907 socket.flush().await.unwrap();
5908 });
5909 }
5910 });
5911
5912 let client = test_http_client();
5913 let cap = 64;
5914
5915 let over = client
5916 .get(format!("http://{addr}/over"))
5917 .send()
5918 .await
5919 .unwrap();
5920 assert_eq!(
5921 over.content_length(),
5922 None,
5923 "chunked response must not declare a length for this test to be meaningful"
5924 );
5925 let err = streamable_http::read_body_capped(over, cap)
5926 .await
5927 .expect_err("a chunked body past the cap must fail, not OOM");
5928 assert!(
5929 err.to_string().contains("exceeds"),
5930 "unexpected error: {err}"
5931 );
5932
5933 let under = client
5934 .get(format!("http://{addr}/under"))
5935 .send()
5936 .await
5937 .unwrap();
5938 let body = streamable_http::read_body_capped(under, cap)
5939 .await
5940 .expect("a chunked body under the cap reads fine");
5941 assert_eq!(body, "x".repeat(16));
5942
5943 server.abort();
5944 }
5945
5946 #[tokio::test]
5947 async fn error_body_excerpt_stops_at_cap_without_waiting_for_eof() {
5948 use tokio::io::{AsyncReadExt, AsyncWriteExt};
5949 use tokio::net::TcpListener;
5950
5951 let _lock = lock_mcp_loopback_tests().await;
5952 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
5953 let addr = listener.local_addr().unwrap();
5954 let server_cancel = tokio_util::sync::CancellationToken::new();
5955 let task_cancel = server_cancel.clone();
5956
5957 // Deliberately omit the terminating zero-sized chunk and keep the socket
5958 // open. A `.text()`-based diagnostic would wait for EOF; the bounded
5959 // reader must return as soon as the first chunk reaches the cap.
5960 let server = tokio::spawn(async move {
5961 let (mut socket, _) = listener.accept().await.unwrap();
5962 let mut request = Vec::new();
5963 let mut buf = [0; 1024];
5964 loop {
5965 let n = socket.read(&mut buf).await.unwrap();
5966 if n == 0 {
5967 return;
5968 }
5969 request.extend_from_slice(&buf[..n]);
5970 if request.windows(4).any(|window| window == b"\r\n\r\n") {
5971 break;
5972 }
5973 }
5974 socket
5975 .write_all(
5976 b"HTTP/1.1 500 Internal Server Error\r\nContent-Type: text/plain\r\nTransfer-Encoding: chunked\r\n\r\n100\r\n",
5977 )
5978 .await
5979 .unwrap();
5980 socket.write_all(&[b'x'; 256]).await.unwrap();
5981 socket.write_all(b"\r\n").await.unwrap();
5982 socket.flush().await.unwrap();
5983 task_cancel.cancelled().await;
5984 });
5985
5986 let response = test_http_client()
5987 .get(format!("http://{addr}/preview"))
5988 .send()
5989 .await
5990 .unwrap();
5991 let preview = tokio::time::timeout(Duration::from_secs(1), bounded_body_excerpt(response, 64))
5992 .await
5993 .expect("bounded excerpt must not wait for an attacker-controlled EOF");
5994 assert_eq!(preview, format!("{}…", "x".repeat(64)));
5995
5996 server_cancel.cancel();
5997 server.abort();
5998 }
5999
6000 #[tokio::test]
6001 async fn streamable_http_stale_session_reconnects_and_retries_tool_call() {
6002 use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
6003 use tokio::io::{AsyncReadExt, AsyncWriteExt};
6004 use tokio::net::TcpListener;
6005
6006 async fn write_response(socket: &mut tokio::net::TcpStream, response: &[u8]) {
6007 socket.write_all(response).await.unwrap();
6008 socket.flush().await.unwrap();
6009 socket.shutdown().await.unwrap();
6010 }
6011
6012 let _lock = lock_mcp_loopback_tests().await;
6013 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6014 let addr = listener.local_addr().unwrap();
6015 let get_count = Arc::new(AtomicUsize::new(0));
6016 let stale_seen = Arc::new(AtomicBool::new(false));
6017 let success_seen = Arc::new(AtomicBool::new(false));
6018 let server_get_count = Arc::clone(&get_count);
6019 let server_stale_seen = Arc::clone(&stale_seen);
6020 let server_success_seen = Arc::clone(&success_seen);
6021
6022 let server = tokio::spawn(async move {
6023 loop {
6024 let Ok((mut socket, _)) = listener.accept().await else {
6025 break;
6026 };
6027 let get_count = Arc::clone(&server_get_count);
6028 let stale_seen = Arc::clone(&server_stale_seen);
6029 let success_seen = Arc::clone(&server_success_seen);
6030 tokio::spawn(async move {
6031 let mut request = Vec::new();
6032 let mut buf = [0; 4096];
6033 let header_end = loop {
6034 let n = socket.read(&mut buf).await.unwrap();
6035 if n == 0 {
6036 return;
6037 }
6038 request.extend_from_slice(&buf[..n]);
6039 if let Some(pos) = request.windows(4).position(|w| w == b"\r\n\r\n") {
6040 break pos + 4;
6041 }
6042 };
6043 let headers = String::from_utf8_lossy(&request[..header_end]).to_string();
6044 let content_length = headers
6045 .lines()
6046 .find_map(|line| {
6047 let (name, value) = line.split_once(':')?;
6048 name.eq_ignore_ascii_case("content-length")
6049 .then(|| value.trim().parse::<usize>().ok())
6050 .flatten()
6051 })
6052 .unwrap_or(0);
6053 while request.len() < header_end + content_length {
6054 let n = socket.read(&mut buf).await.unwrap();
6055 if n == 0 {
6056 return;
6057 }
6058 request.extend_from_slice(&buf[..n]);
6059 }
6060 let body = &request[header_end..header_end + content_length];
6061 let session_header = headers.lines().find_map(|line| {
6062 let (name, value) = line.split_once(':')?;
6063 name.eq_ignore_ascii_case("mcp-session-id")
6064 .then(|| value.trim().to_string())
6065 });
6066
6067 if headers.starts_with("GET /mcp ") {
6068 let count = get_count.fetch_add(1, AtomicOrdering::SeqCst);
6069 let session = if count == 0 { "sess-old" } else { "sess-new" };
6070 let response = format!(
6071 "HTTP/1.1 200 OK\r\nConnection: close\r\nMcp-Session-Id: {session}\r\nContent-Length: 0\r\n\r\n"
6072 );
6073 write_response(&mut socket, response.as_bytes()).await;
6074 return;
6075 }
6076
6077 let request_json: serde_json::Value = serde_json::from_slice(body).unwrap();
6078 let method = request_json
6079 .get("method")
6080 .and_then(serde_json::Value::as_str)
6081 .unwrap_or("");
6082 let id = request_json
6083 .get("id")
6084 .cloned()
6085 .unwrap_or_else(|| serde_json::json!("0"));
6086
6087 if method == "tools/call" && session_header.as_deref() == Some("sess-old") {
6088 stale_seen.store(true, AtomicOrdering::SeqCst);
6089 write_response(
6090 &mut socket,
6091 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\"}",
6092 )
6093 .await;
6094 return;
6095 }
6096
6097 let result = match method {
6098 "initialize" => serde_json::json!({
6099 "protocolVersion": "2024-11-05",
6100 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
6101 }),
6102 "tools/list" => serde_json::json!({
6103 "tools": [
6104 { "name": "search", "inputSchema": {} }
6105 ]
6106 }),
6107 "resources/list" => serde_json::json!({ "resources": [] }),
6108 "resources/templates/list" => {
6109 serde_json::json!({ "resourceTemplates": [] })
6110 }
6111 "prompts/list" => serde_json::json!({ "prompts": [] }),
6112 "tools/call" => {
6113 assert_eq!(session_header.as_deref(), Some("sess-new"));
6114 success_seen.store(true, AtomicOrdering::SeqCst);
6115 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
6116 }
6117 _ => {
6118 write_response(
6119 &mut socket,
6120 b"HTTP/1.1 202 Accepted\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
6121 )
6122 .await;
6123 return;
6124 }
6125 };
6126 let response_body = serde_json::json!({
6127 "jsonrpc": "2.0",
6128 "id": id,
6129 "result": result
6130 })
6131 .to_string();
6132 let response = format!(
6133 "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}",
6134 response_body.len(),
6135 response_body
6136 );
6137 write_response(&mut socket, response.as_bytes()).await;
6138 });
6139 }
6140 });
6141
6142 let mut cfg = McpConfig::default();
6143 cfg.servers.insert(
6144 "dephy".to_string(),
6145 McpServerConfig {
6146 command: None,
6147 args: Vec::new(),
6148 env: HashMap::new(),
6149 cwd: None,
6150 url: Some(format!("http://{addr}/mcp")),
6151 transport: None,
6152 connect_timeout: Some(10),
6153 execute_timeout: Some(10),
6154 read_timeout: None,
6155 disabled: false,
6156 enabled: true,
6157 required: false,
6158 enabled_tools: Vec::new(),
6159 disabled_tools: Vec::new(),
6160 headers: HashMap::new(),
6161 env_headers: HashMap::new(),
6162 bearer_token_env_var: None,
6163 scopes: Vec::new(),
6164 oauth: None,
6165 oauth_resource: None,
6166 reviewed_plugin: None,
6167 runtime_added: false,
6168 allow_private_network: false,
6169 },
6170 );
6171 let mut pool = McpPool::new(cfg);
6172
6173 let result = pool
6174 .call_tool("mcp_dephy_search", serde_json::json!({ "query": "dephy" }))
6175 .await
6176 .unwrap();
6177
6178 assert_eq!(
6179 result,
6180 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
6181 );
6182 assert!(stale_seen.load(AtomicOrdering::SeqCst));
6183 assert!(success_seen.load(AtomicOrdering::SeqCst));
6184 assert_eq!(get_count.load(AtomicOrdering::SeqCst), 2);
6185
6186 server.abort();
6187 }
6188
6189 #[tokio::test]
6190 async fn legacy_sse_session_expiry_is_marked_stale() {
6191 use tokio::io::{AsyncReadExt, AsyncWriteExt};
6192 use tokio::net::TcpListener;
6193 use tokio::sync::mpsc;
6194
6195 let _lock = lock_mcp_loopback_tests().await;
6196 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6197 let addr = listener.local_addr().unwrap();
6198
6199 let server = tokio::spawn(async move {
6200 let (mut socket, _) = listener.accept().await.unwrap();
6201 let mut request = Vec::new();
6202 let mut buf = [0; 4096];
6203 let header_end = loop {
6204 let n = socket.read(&mut buf).await.unwrap();
6205 if n == 0 {
6206 return;
6207 }
6208 request.extend_from_slice(&buf[..n]);
6209 if let Some(pos) = request.windows(4).position(|w| w == b"\r\n\r\n") {
6210 break pos + 4;
6211 }
6212 };
6213 let headers = String::from_utf8_lossy(&request[..header_end]);
6214 assert!(headers.starts_with("POST /messages "));
6215 socket
6216 .write_all(
6217 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\"}",
6218 )
6219 .await
6220 .unwrap();
6221 });
6222
6223 let (_sender, receiver) = mpsc::channel(1);
6224 let sse_task = tokio::spawn(async {});
6225 let mut transport = SseTransport {
6226 client: test_mcp_http_client(&format!("http://{addr}/sse")),
6227 base_url: format!("http://{addr}/sse"),
6228 endpoint_url: Some(format!("http://{addr}/messages")),
6229 receiver,
6230 sse_task,
6231 };
6232
6233 let err = transport
6234 .send(br#"{"jsonrpc":"2.0","id":1,"method":"tools/call"}"#.to_vec())
6235 .await
6236 .expect_err("expired SSE session should fail");
6237
6238 assert!(
6239 is_mcp_stale_session_error(&err),
6240 "SSE session expiry should be retryable, got: {err:#}"
6241 );
6242
6243 server.abort();
6244 }
6245
6246 /// Read one HTTP/1.1 request from a legacy SSE test server socket.
6247 async fn read_legacy_sse_http_request(
6248 socket: &mut tokio::net::TcpStream,
6249 ) -> (String, serde_json::Value) {
6250 let mut request = Vec::new();
6251 let mut buf = [0; 4096];
6252 let header_end = loop {
6253 let n = tokio::io::AsyncReadExt::read(socket, &mut buf)
6254 .await
6255 .unwrap();
6256 if n == 0 {
6257 return (String::new(), serde_json::Value::Null);
6258 }
6259 request.extend_from_slice(&buf[..n]);
6260 if let Some(pos) = request.windows(4).position(|w| w == b"\r\n\r\n") {
6261 break pos + 4;
6262 }
6263 };
6264 let headers = String::from_utf8_lossy(&request[..header_end]).to_string();
6265 let content_length = headers
6266 .lines()
6267 .find_map(|line| {
6268 let (name, value) = line.split_once(':')?;
6269 name.eq_ignore_ascii_case("content-length")
6270 .then(|| value.trim().parse::<usize>().ok())
6271 .flatten()
6272 })
6273 .unwrap_or(0);
6274 while request.len() < header_end + content_length {
6275 let n = tokio::io::AsyncReadExt::read(socket, &mut buf)
6276 .await
6277 .unwrap();
6278 if n == 0 {
6279 return (headers, serde_json::Value::Null);
6280 }
6281 request.extend_from_slice(&buf[..n]);
6282 }
6283 let body = &request[header_end..header_end + content_length];
6284 let json = if body.is_empty() {
6285 serde_json::Value::Null
6286 } else {
6287 serde_json::from_slice(body).unwrap()
6288 };
6289 (headers, json)
6290 }
6291
6292 #[tokio::test]
6293 async fn legacy_sse_closed_stream_reports_unknown_outcome_and_reconnects_without_replay() {
6294 use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
6295 use tokio::io::AsyncWriteExt;
6296 use tokio::net::TcpListener;
6297 use tokio::sync::mpsc;
6298
6299 // A concurrent proxy fixture changes process-wide HTTP_PROXY/NO_PROXY.
6300 // Hold the environment guard before the loopback guard, as other MCP
6301 // tests do, so this server is always reached directly.
6302 let _env = crate::test_support::lock_test_env();
6303 let _lock = lock_mcp_loopback_tests().await;
6304 let _no_proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
6305 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6306 let addr = listener.local_addr().unwrap();
6307 let active_sse = Arc::new(Mutex::new(None::<mpsc::UnboundedSender<Option<String>>>));
6308 let get_count = Arc::new(AtomicUsize::new(0));
6309 let tool_call_count = Arc::new(AtomicUsize::new(0));
6310 let success_seen = Arc::new(AtomicBool::new(false));
6311 let server_active_sse = Arc::clone(&active_sse);
6312 let server_get_count = Arc::clone(&get_count);
6313 let server_tool_call_count = Arc::clone(&tool_call_count);
6314 let server_success_seen = Arc::clone(&success_seen);
6315
6316 let server = tokio::spawn(async move {
6317 loop {
6318 let Ok((mut socket, _)) = listener.accept().await else {
6319 break;
6320 };
6321 let active_sse = Arc::clone(&server_active_sse);
6322 let get_count = Arc::clone(&server_get_count);
6323 let tool_call_count = Arc::clone(&server_tool_call_count);
6324 let success_seen = Arc::clone(&server_success_seen);
6325 tokio::spawn(async move {
6326 let (headers, request_json) = read_legacy_sse_http_request(&mut socket).await;
6327 if headers.starts_with("GET /sse ") {
6328 get_count.fetch_add(1, AtomicOrdering::SeqCst);
6329 let (tx, mut rx) = mpsc::unbounded_channel::<Option<String>>();
6330 *active_sse.lock().unwrap() = Some(tx);
6331 socket
6332 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
6333 .await
6334 .unwrap();
6335 socket
6336 .write_all(b"event: endpoint\ndata: /messages\n\n")
6337 .await
6338 .unwrap();
6339 while let Some(message) = rx.recv().await {
6340 let Some(message) = message else {
6341 return;
6342 };
6343 let event = format!("event: message\ndata: {message}\n\n");
6344 socket.write_all(event.as_bytes()).await.unwrap();
6345 }
6346 return;
6347 }
6348
6349 if !headers.starts_with("POST /messages ") {
6350 return;
6351 }
6352
6353 socket
6354 .write_all(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
6355 .await
6356 .unwrap();
6357
6358 let method = request_json
6359 .get("method")
6360 .and_then(serde_json::Value::as_str)
6361 .unwrap_or("");
6362 if method == "notifications/initialized" {
6363 return;
6364 }
6365
6366 let id = request_json
6367 .get("id")
6368 .cloned()
6369 .unwrap_or_else(|| serde_json::json!("0"));
6370
6371 if method == "tools/call" {
6372 let count = tool_call_count.fetch_add(1, AtomicOrdering::SeqCst);
6373 if count == 0 {
6374 if let Some(tx) = active_sse.lock().unwrap().take() {
6375 let _ = tx.send(None);
6376 }
6377 return;
6378 }
6379 }
6380
6381 let result = match method {
6382 "initialize" => serde_json::json!({
6383 "protocolVersion": "2024-11-05",
6384 "capabilities": {"tools": {}, "resources": {}, "prompts": {}}
6385 }),
6386 "tools/list" => serde_json::json!({
6387 "tools": [
6388 { "name": "search", "inputSchema": {} }
6389 ]
6390 }),
6391 "resources/list" => serde_json::json!({ "resources": [] }),
6392 "resources/templates/list" => {
6393 serde_json::json!({ "resourceTemplates": [] })
6394 }
6395 "prompts/list" => serde_json::json!({ "prompts": [] }),
6396 "tools/call" => {
6397 success_seen.store(true, AtomicOrdering::SeqCst);
6398 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
6399 }
6400 other => panic!("unexpected method: {other}"),
6401 };
6402 let response = serde_json::json!({
6403 "jsonrpc": "2.0",
6404 "id": id,
6405 "result": result
6406 })
6407 .to_string();
6408 // Deliver the response over the *current* SSE channel. The
6409 // retry tool call can race ahead of the reconnecting GET
6410 // /sse that re-stores the sender; under parallel load those
6411 // two server tasks are scheduled in either order, so wait
6412 // briefly for the channel instead of dropping the response
6413 // (which left the client hanging until timeout) (#2597).
6414 let send_deadline = std::time::Instant::now() + std::time::Duration::from_secs(5);
6415 let tx = loop {
6416 if let Some(tx) = active_sse.lock().unwrap().as_ref().cloned() {
6417 break Some(tx);
6418 }
6419 if std::time::Instant::now() >= send_deadline {
6420 break None;
6421 }
6422 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
6423 };
6424 if let Some(tx) = tx {
6425 let _ = tx.send(Some(response));
6426 }
6427 });
6428 }
6429 });
6430
6431 let mut cfg = McpConfig::default();
6432 cfg.servers.insert(
6433 "dephy".to_string(),
6434 McpServerConfig {
6435 command: None,
6436 args: Vec::new(),
6437 env: HashMap::new(),
6438 cwd: None,
6439 url: Some(format!("http://{addr}/sse")),
6440 transport: Some("sse".to_string()),
6441 connect_timeout: Some(10),
6442 execute_timeout: Some(10),
6443 read_timeout: None,
6444 disabled: false,
6445 enabled: true,
6446 required: false,
6447 enabled_tools: Vec::new(),
6448 disabled_tools: Vec::new(),
6449 headers: HashMap::new(),
6450 env_headers: HashMap::new(),
6451 bearer_token_env_var: None,
6452 scopes: Vec::new(),
6453 oauth: None,
6454 oauth_resource: None,
6455 reviewed_plugin: None,
6456 runtime_added: false,
6457 allow_private_network: false,
6458 },
6459 );
6460 let mut pool = McpPool::new(cfg);
6461
6462 // The server received the call and then closed the stream: it may have
6463 // run the tool, so the call must fail as unknown instead of replaying.
6464 let err = pool
6465 .call_tool("mcp_dephy_search", serde_json::json!({ "query": "dephy" }))
6466 .await
6467 .expect_err("a call whose transport closed mid-flight must not be replayed");
6468 assert!(
6469 format!("{err:#}").contains("outcome unknown, not retried"),
6470 "unexpected error: {err:#}"
6471 );
6472 assert_eq!(tool_call_count.load(AtomicOrdering::SeqCst), 1);
6473 assert!(!success_seen.load(AtomicOrdering::SeqCst));
6474
6475 // The dead connection was dropped, so the next call reconnects.
6476 let result = pool
6477 .call_tool("mcp_dephy_search", serde_json::json!({ "query": "dephy" }))
6478 .await
6479 .unwrap();
6480 assert_eq!(
6481 result,
6482 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
6483 );
6484 assert_eq!(tool_call_count.load(AtomicOrdering::SeqCst), 2);
6485 assert_eq!(get_count.load(AtomicOrdering::SeqCst), 2);
6486 assert!(success_seen.load(AtomicOrdering::SeqCst));
6487
6488 server.abort();
6489 }
6490
6491 /// Servers without an explicit `transport = "sse"` reach legacy SSE through
6492 /// the Streamable HTTP fallback, inside `HttpTransport`. An event stream that
6493 /// closes while idle must read as not ready there too, so the next call
6494 /// reconnects before it dispatches instead of losing the result.
6495 #[tokio::test]
6496 async fn fallback_sse_stream_closed_while_idle_reconnects_before_dispatch() {
6497 use tokio::io::AsyncWriteExt;
6498 use tokio::net::TcpListener;
6499 use tokio::sync::mpsc;
6500
6501 let _env = crate::test_support::lock_test_env();
6502 let _lock = lock_mcp_loopback_tests().await;
6503 let _no_proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
6504 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6505 let addr = listener.local_addr().unwrap();
6506 let active_sse = Arc::new(Mutex::new(None::<mpsc::UnboundedSender<Option<String>>>));
6507 let initialize_count = Arc::new(AtomicUsize::new(0));
6508 let tool_call_count = Arc::new(AtomicUsize::new(0));
6509 let server_active_sse = Arc::clone(&active_sse);
6510 let server_initialize_count = Arc::clone(&initialize_count);
6511 let server_tool_call_count = Arc::clone(&tool_call_count);
6512
6513 let server = tokio::spawn(async move {
6514 while let Ok((mut socket, _)) = listener.accept().await {
6515 let active_sse = Arc::clone(&server_active_sse);
6516 let initialize_count = Arc::clone(&server_initialize_count);
6517 let tool_call_count = Arc::clone(&server_tool_call_count);
6518 tokio::spawn(async move {
6519 let (headers, request_json) = read_legacy_sse_http_request(&mut socket).await;
6520 if headers.starts_with("GET /mcp ") {
6521 // The session preflight and the fallback stream both land
6522 // here; the later GET replaces the earlier stream.
6523 let (tx, mut rx) = mpsc::unbounded_channel::<Option<String>>();
6524 *active_sse.lock().unwrap() = Some(tx);
6525 let opened = socket
6526 .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\nevent: endpoint\ndata: /messages\n\n")
6527 .await;
6528 if opened.is_err() {
6529 return;
6530 }
6531 while let Some(Some(message)) = rx.recv().await {
6532 let event = format!("event: message\ndata: {message}\n\n");
6533 if socket.write_all(event.as_bytes()).await.is_err() {
6534 return;
6535 }
6536 }
6537 return;
6538 }
6539 if headers.starts_with("POST /mcp ") {
6540 // No Streamable HTTP: the client falls back to legacy SSE.
6541 let _ = socket
6542 .write_all(b"HTTP/1.1 405 Method Not Allowed\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
6543 .await;
6544 return;
6545 }
6546 if !headers.starts_with("POST /messages ") {
6547 return;
6548 }
6549 let _ = socket
6550 .write_all(
6551 b"HTTP/1.1 202 Accepted\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
6552 )
6553 .await;
6554 let method = request_json
6555 .get("method")
6556 .and_then(serde_json::Value::as_str)
6557 .unwrap_or("");
6558 let result = match method {
6559 "notifications/initialized" => return,
6560 "initialize" => {
6561 initialize_count.fetch_add(1, AtomicOrdering::SeqCst);
6562 serde_json::json!({
6563 "protocolVersion": "2024-11-05",
6564 "capabilities": {"tools": {}}
6565 })
6566 }
6567 "tools/list" => serde_json::json!({
6568 "tools": [{ "name": "search", "inputSchema": {} }]
6569 }),
6570 "resources/list" => serde_json::json!({ "resources": [] }),
6571 "resources/templates/list" => serde_json::json!({ "resourceTemplates": [] }),
6572 "prompts/list" => serde_json::json!({ "prompts": [] }),
6573 "tools/call" => {
6574 tool_call_count.fetch_add(1, AtomicOrdering::SeqCst);
6575 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
6576 }
6577 other => panic!("unexpected method: {other}"),
6578 };
6579 let response = serde_json::json!({
6580 "jsonrpc": "2.0",
6581 "id": request_json.get("id").cloned().unwrap_or(serde_json::Value::Null),
6582 "result": result
6583 })
6584 .to_string();
6585 let tx = active_sse.lock().unwrap().as_ref().cloned();
6586 if let Some(tx) = tx {
6587 let _ = tx.send(Some(response));
6588 }
6589 });
6590 }
6591 });
6592
6593 let mut cfg = McpConfig::default();
6594 let mut server_config = test_server_config();
6595 server_config.command = None;
6596 server_config.url = Some(format!("http://{addr}/mcp"));
6597 server_config.connect_timeout = Some(10);
6598 server_config.execute_timeout = Some(10);
6599 cfg.servers.insert("fallback".to_string(), server_config);
6600 let mut pool = McpPool::new(cfg);
6601 let ok = serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] });
6602
6603 let result = pool
6604 .call_tool("mcp_fallback_search", serde_json::json!({}))
6605 .await
6606 .unwrap();
6607 assert_eq!(result, ok);
6608
6609 // The server ends the event stream while the connection is idle.
6610 let stream = active_sse
6611 .lock()
6612 .unwrap()
6613 .take()
6614 .expect("fallback stream open");
6615 stream.send(None).unwrap();
6616 let deadline = tokio::time::Instant::now() + Duration::from_secs(5);
6617 while pool.connections["fallback"].is_transport_ready() {
6618 assert!(
6619 tokio::time::Instant::now() < deadline,
6620 "a closed fallback SSE stream must stop reading as ready"
6621 );
6622 tokio::time::sleep(Duration::from_millis(20)).await;
6623 }
6624
6625 // The next call reconnects first, so it is dispatched once and answered.
6626 let result = pool
6627 .call_tool("mcp_fallback_search", serde_json::json!({}))
6628 .await
6629 .unwrap();
6630 assert_eq!(result, ok);
6631 assert_eq!(tool_call_count.load(AtomicOrdering::SeqCst), 2);
6632 assert_eq!(initialize_count.load(AtomicOrdering::SeqCst), 2);
6633
6634 server.abort();
6635 }
6636
6637 #[test]
6638 fn session_id_starts_none() {
6639 let transport = StreamableHttpTransport::new(
6640 test_mcp_http_client("https://example.invalid/mcp"),
6641 "https://example.invalid/mcp".to_string(),
6642 );
6643 assert!(transport.session_id.is_none());
6644 }
6645
6646 /// Session ID captured from a POST response is replayed on the next POST.
6647 #[tokio::test]
6648 async fn session_id_captured_from_post_response_and_replayed() {
6649 use tokio::io::{AsyncReadExt, AsyncWriteExt};
6650 use tokio::net::TcpListener;
6651
6652 let _lock = lock_mcp_loopback_tests().await;
6653 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6654 let addr = listener.local_addr().unwrap();
6655 let server = tokio::spawn(async move {
6656 let (mut socket, _) = listener.accept().await.unwrap();
6657 let mut buf = [0u8; 4096];
6658 let n = socket.read(&mut buf).await.unwrap();
6659 let req = String::from_utf8_lossy(&buf[..n]);
6660 assert!(req.starts_with("POST "), "expected POST, got: {req}");
6661
6662 // First POST: return a session ID so the transport captures it.
6663 socket
6664 .write_all(
6665 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{}",
6666 )
6667 .await
6668 .unwrap();
6669 socket.flush().await.unwrap();
6670
6671 // Read the second POST — should contain the session ID.
6672 let mut buf2 = [0u8; 4096];
6673 let n2 = socket.read(&mut buf2).await.unwrap();
6674 let req2 = String::from_utf8_lossy(&buf2[..n2]);
6675 // reqwest lower-cases header names.
6676 let req2_lower = req2.to_lowercase();
6677 assert!(
6678 req2_lower.contains("mcp-session-id: sess-abc-123"),
6679 "second POST must replay captured session ID, got:\n{req2}"
6680 );
6681
6682 socket
6683 .write_all(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
6684 .await
6685 .unwrap();
6686 });
6687
6688 let url = format!("http://{addr}/mcp");
6689 let mut transport = StreamableHttpTransport::new(test_mcp_http_client(&url), url);
6690
6691 // First send: server returns Mcp-Session-Id.
6692 transport
6693 .send(json_frame(serde_json::json!({
6694 "jsonrpc": "2.0", "id": 1,
6695 "method": "initialize",
6696 "params": {}
6697 })))
6698 .await
6699 .unwrap();
6700 assert_eq!(
6701 transport.session_id.as_deref(),
6702 Some("sess-abc-123"),
6703 "session ID should be captured from response"
6704 );
6705
6706 // Second send: should replay the session ID.
6707 transport
6708 .send(json_frame(serde_json::json!({
6709 "jsonrpc": "2.0", "id": 2,
6710 "method": "tools/list",
6711 "params": {}
6712 })))
6713 .await
6714 .unwrap();
6715
6716 server.abort();
6717 }
6718
6719 /// Custom headers configured in McpServerConfig are applied to the GET
6720 /// preflight so servers that require auth on session-establishment GET
6721 /// (e.g. Hindsight, #1629) can authenticate it.
6722 #[tokio::test]
6723 async fn custom_headers_applied_to_get_preflight() {
6724 use tokio::io::{AsyncReadExt, AsyncWriteExt};
6725 use tokio::net::TcpListener;
6726
6727 // Lock order is env first, then loopback — the OAuth pool tests take
6728 // them in that order, and inverting them deadlocks the suite.
6729 let _env = crate::test_support::lock_test_env();
6730 let _lock = lock_mcp_loopback_tests().await;
6731 // The fixture client honors an operator-configured proxy; pin loopback
6732 // out of any ambient proxy so a concurrent proxy-configuring test can
6733 // never route this GET away from the fixture server.
6734 let _no_proxy = crate::test_support::EnvVarGuard::set("NO_PROXY", "*");
6735 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
6736 let addr = listener.local_addr().unwrap();
6737 // The test signals success by writing to this flag — the GET handler
6738 // sets it when it sees the expected header.
6739 let header_seen = Arc::new(AtomicBool::new(false));
6740 let header_seen_srv = Arc::clone(&header_seen);
6741
6742 let server = tokio::spawn(async move {
6743 let (mut socket, _) = listener.accept().await.unwrap();
6744 let mut buf = [0u8; 4096];
6745 let n = socket.read(&mut buf).await.unwrap();
6746 let req = String::from_utf8_lossy(&buf[..n]);
6747
6748 // reqwest lower-cases header names.
6749 if req.starts_with("GET ") && req.to_lowercase().contains("x-custom-auth: my-test-token") {
6750 header_seen_srv.store(true, AtomicOrdering::SeqCst);
6751 }
6752
6753 socket
6754 .write_all(b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
6755 .await
6756 .unwrap();
6757 });
6758
6759 let url = format!("http://{addr}/mcp");
6760 let mut headers = HashMap::new();
6761 headers.insert("X-Custom-Auth".to_string(), "my-test-token".to_string());
6762
6763 let mut transport = HttpTransport::new(
6764 test_mcp_http_client(&url).with_mcp_auth(McpHttpAuth {
6765 headers,
6766 ..Default::default()
6767 }),
6768 url,
6769 tokio_util::sync::CancellationToken::new(),
6770 Duration::from_secs(10),
6771 );
6772
6773 transport.try_establish_session().await.unwrap();
6774
6775 server.abort();
6776
6777 assert!(
6778 header_seen.load(AtomicOrdering::SeqCst),
6779 "GET preflight must include user-configured custom headers"
6780 );
6781 }
6782
6783 // === add_runtime_server_config conflict tests ===
6784
6785 #[test]
6786 fn add_runtime_server_config_rejects_static_conflict() {
6787 let config: McpConfig = serde_json::from_str(
6788 r#"{
6789 "servers": {
6790 "existing": {"command": "node server.js"}
6791 }
6792 }"#,
6793 )
6794 .unwrap();
6795 let pool = McpPool::new(config);
6796
6797 let err = pool
6798 .add_runtime_server_config(
6799 "existing".to_string(),
6800 serde_json::from_str(r#"{"command": "npx other"}"#).unwrap(),
6801 )
6802 .unwrap_err();
6803 assert!(err.contains("already exists in the config file"));
6804 }
6805
6806 #[test]
6807 fn add_runtime_server_config_rejects_dynamic_duplicate() {
6808 let pool = McpPool::new(McpConfig::default());
6809
6810 pool.add_runtime_server_config(
6811 "my_server".to_string(),
6812 serde_json::from_str(r#"{"command": "node a.js"}"#).unwrap(),
6813 )
6814 .unwrap();
6815
6816 let err = pool
6817 .add_runtime_server_config(
6818 "my_server".to_string(),
6819 serde_json::from_str(r#"{"command": "node b.js"}"#).unwrap(),
6820 )
6821 .unwrap_err();
6822 assert!(err.contains("already started earlier"));
6823 }
6824
6825 #[test]
6826 fn add_runtime_server_config_accepts_new_name() {
6827 let pool = McpPool::new(McpConfig::default());
6828
6829 pool.add_runtime_server_config(
6830 "brand_new".to_string(),
6831 serde_json::from_str(r#"{"command": "node x.js"}"#).unwrap(),
6832 )
6833 .unwrap();
6834 }
6835
6836 /// Server attribution and the model-facing tool name must come from one
6837 /// definition. If they ever drift, a human reading tool provenance would be
6838 /// told which server owns a name the model never saw.
6839 #[test]
6840 fn mcp_model_tool_names_and_server_attribution_share_one_definition() {
6841 assert_eq!(
6842 McpPool::mcp_model_tool_name("files", "read"),
6843 "mcp_files_read"
6844 );
6845 // A server name containing `_` is exactly why the reverse split is a guess.
6846 assert_eq!(
6847 McpPool::mcp_model_tool_name("my_server", "read_file"),
6848 "mcp_my_server_read_file"
6849 );
6850
6851 let resolved =
6852 McpPool::resolve_tool_server_map([("files", "read"), ("git", "status")].into_iter());
6853 assert_eq!(
6854 resolved.get("mcp_files_read").map(String::as_str),
6855 Some("files")
6856 );
6857 assert_eq!(
6858 resolved.get("mcp_git_status").map(String::as_str),
6859 Some("git")
6860 );
6861
6862 // Ambiguity: two servers collapse onto the same model name. Neither wins,
6863 // so the name resolves to no server and callers report it as unknown —
6864 // the same rule `all_tools` applies when it hides the ambiguous tool.
6865 let ambiguous = McpPool::resolve_tool_server_map(
6866 [("a_b", "c"), ("a", "b_c"), ("solo", "tool")].into_iter(),
6867 );
6868 assert!(
6869 !ambiguous.contains_key("mcp_a_b_c"),
6870 "an ambiguous model name must resolve to no server"
6871 );
6872 assert_eq!(
6873 ambiguous.get("mcp_solo_tool").map(String::as_str),
6874 Some("solo")
6875 );
6876 }
6877
6878 #[test]
6879 fn removed_runtime_server_config_can_be_retried_with_same_name() {
6880 let mut pool = McpPool::new(McpConfig::default());
6881 let config: McpServerConfig = serde_json::from_str(r#"{"command": "node a.js"}"#).unwrap();
6882
6883 pool.add_runtime_server_config("retryable".to_string(), config.clone())
6884 .unwrap();
6885 pool.remove_runtime_server_config("retryable");
6886 pool.add_runtime_server_config("retryable".to_string(), config)
6887 .expect("rollback must release the deterministic runtime name");
6888 }
6889
6890 #[tokio::test]
6891 async fn mcp_initialize_sends_empty_client_capabilities_and_accepts_2025_11_25() {
6892 let sent = Arc::new(Mutex::new(Vec::new()));
6893 let transport = ScriptedValueTransport {
6894 sent: Arc::clone(&sent),
6895 responses: VecDeque::from([json_frame(serde_json::json!({
6896 "jsonrpc": "2.0",
6897 "id": 1,
6898 "result": {
6899 "protocolVersion": "2025-11-25",
6900 "serverInfo": {"name": "current-sdk", "version": "1.0.0"},
6901 "capabilities": {"tools": {}}
6902 }
6903 }))]),
6904 };
6905 let mut conn = test_connection(Box::new(transport));
6906
6907 conn.initialize()
6908 .await
6909 .expect("a 2025-11-25 server must complete the handshake");
6910
6911 let sent = sent.lock().unwrap();
6912 let initialize = sent
6913 .iter()
6914 .find(|message| message["method"] == "initialize")
6915 .expect("initialize sent");
6916 assert_eq!(
6917 initialize["params"],
6918 serde_json::json!({
6919 "protocolVersion": MCP_PROTOCOL_VERSION,
6920 "clientInfo": {"name": "codewhale-tui", "version": env!("CARGO_PKG_VERSION")},
6921 "capabilities": {}
6922 }),
6923 "client capabilities must not declare server-side tools/resources/prompts"
6924 );
6925 }
6926
6927 #[tokio::test]
6928 async fn mcp_initialize_rejection_from_expired_aws_sso_names_the_login_command() {
6929 let transport = ScriptedValueTransport {
6930 sent: Arc::new(Mutex::new(Vec::new())),
6931 responses: VecDeque::from([json_frame(serde_json::json!({
6932 "jsonrpc": "2.0",
6933 "id": 1,
6934 "error": {
6935 "code": -32602,
6936 "message": "Error retrieving credentials: The SSO session associated with this profile has expired or is otherwise invalid."
6937 }
6938 }))]),
6939 };
6940 let mut conn = test_connection(Box::new(transport));
6941 conn.config.args = vec![
6942 "mcp-proxy-for-aws@1.6.4".to_string(),
6943 "--profile".to_string(),
6944 "work-sso".to_string(),
6945 ];
6946
6947 let error = format!("{:#}", conn.initialize().await.expect_err("rejected"));
6948 assert!(
6949 error.contains(
6950 "run `aws sso login --profile work-sso` in a terminal, then `/mcp retry mock`"
6951 ),
6952 "{error}"
6953 );
6954 assert_eq!(
6955 mcp_recovery_kind(true, true, false, Some(&error), false),
6956 Some(McpRecoveryKind::AwsLogin)
6957 );
6958 }
6959
6960 #[test]
6961 fn mcp_recovery_kind_routes_expired_aws_credentials_to_external_login() {
6962 for stderr in [
6963 "The SSO session associated with this profile has expired or is otherwise invalid. To refresh this SSO session run aws sso login with the corresponding profile.",
6964 "Error loading SSO Token: Token for my-sso does not exist",
6965 "An error occurred (ExpiredTokenException) when calling the GetCallerIdentity operation: The security token included in the request is expired",
6966 // An expired AWS token that also says 401/Unauthorized must not be
6967 // routed to `/mcp login`, an OAuth flow that cannot renew it.
6968 "UnauthorizedException (401): AWS token has expired",
6969 ] {
6970 assert_eq!(
6971 mcp_recovery_kind(true, true, false, Some(stderr), false),
6972 Some(McpRecoveryKind::AwsLogin),
6973 "{stderr}"
6974 );
6975 }
6976 // Unrelated text that merely shares substrings stays where it was.
6977 for (stderr, expected) in [
6978 ("processor token invalid", McpRecoveryKind::Diagnose),
6979 ("401 Unauthorized", McpRecoveryKind::Reauth),
6980 (
6981 "oauth token expired; invalid_grant",
6982 McpRecoveryKind::Reauth,
6983 ),
6984 (
6985 "laws of token expiry were expired",
6986 McpRecoveryKind::Diagnose,
6987 ),
6988 ] {
6989 assert_eq!(
6990 mcp_recovery_kind(true, true, false, Some(stderr), false),
6991 Some(expected),
6992 "{stderr}"
6993 );
6994 }
6995 assert_eq!(
6996 McpRecoveryKind::AwsLogin.slash_command("aws"),
6997 "/mcp retry aws"
6998 );
6999 assert_eq!(
7000 McpRecoveryKind::AwsLogin.slash_command("name with spaces"),
7001 "/mcp reload"
7002 );
7003
7004 let mut config = test_server_config();
7005 assert_eq!(
7006 aws_login_hint(&config, "aws", ""),
7007 "AWS credentials expired: run `aws sso login` in a terminal, then `/mcp retry aws`"
7008 );
7009 config
7010 .env
7011 .insert("AWS_PROFILE".to_string(), "from-env".to_string());
7012 assert!(aws_login_hint(&config, "aws", "").contains("aws sso login --profile from-env"));
7013 config.args = vec!["--profile=from-arg".to_string()];
7014 assert!(aws_login_hint(&config, "aws", "").contains("aws sso login --profile from-arg"));
7015 // A profile the shell could misread is left out rather than quoted.
7016 config.args = vec!["--profile".to_string(), "x; rm -rf ~".to_string()];
7017 config.env.clear();
7018 assert!(aws_login_hint(&config, "aws", "").contains("run `aws sso login` in"));
7019 config.args = vec!["--profile=--no-sign-request".to_string()];
7020 assert!(aws_login_hint(&config, "aws", "").contains("run `aws sso login` in"));
7021 }
7022
7023 /// Review of #6789: the AWS-before-OAuth ordering lived only in the free
7024 /// classifier, while the panel read the typed needs-auth flag first. A stdio
7025 /// AWS proxy whose token expired with `401`/`Unauthorized` wording must not
7026 /// enter the needs-auth set, must not read `auth_required`, and its row must
7027 /// route to the in-place retry with the login command in the detail.
7028 #[test]
7029 fn expired_aws_login_with_401_wording_never_reaches_oauth_needs_auth() {
7030 let dir = tempfile::tempdir().unwrap();
7031 let mut config = McpConfig::default();
7032 let mut server = test_server_config();
7033 server.args = vec!["--profile".to_string(), "work".to_string()];
7034 config.servers.insert("aws".to_string(), server);
7035 let mut pool = McpPool::new(config);
7036
7037 // Not the initialize branch: no hint was appended to this error.
7038 let error = anyhow::anyhow!("UnauthorizedException (401): AWS token has expired");
7039 pool.note_connect_failure("aws", &error);
7040 assert!(!pool.server_needs_auth("aws"));
7041
7042 let errors = HashMap::from([("aws".to_string(), format_mcp_error_for_display(&error))]);
7043 let snapshot = pool.manager_snapshot(&dir.path().join("mcp.json"), false, &errors);
7044 let aws = snapshot
7045 .servers
7046 .iter()
7047 .find(|server| server.name == "aws")
7048 .expect("aws in snapshot");
7049 assert!(!aws.auth_required, "{aws:?}");
7050 let recovery = aws
7051 .recovery_kind(false)
7052 .expect("a failed server has a recovery");
7053 assert_eq!(recovery, McpRecoveryKind::AwsLogin);
7054 assert_eq!(recovery.slash_command("aws"), "/mcp retry aws");
7055 // Labelled for what it runs, not "re-auth".
7056 assert_eq!(
7057 recovery.label_key(),
7058 codewhale_localization::MessageId::ExtensionsActionReconnect
7059 );
7060 let detail = aws.error.as_deref().unwrap_or_default();
7061 assert!(
7062 detail.contains("run `aws sso login --profile work` in a terminal, then `/mcp retry aws`"),
7063 "{detail}"
7064 );
7065
7066 // A genuine OAuth 401 on the same pool still enters the needs-auth set.
7067 pool.note_connect_failure("aws", &anyhow::anyhow!("HTTP 401 Unauthorized"));
7068 assert!(pool.server_needs_auth("aws"));
7069 }
7070
7071 #[tokio::test]
7072 async fn live_aws_expiry_retires_the_catalog_and_lists_external_recovery() {
7073 for method in ["tools/call", "resources/read", "prompts/get"] {
7074 let dir = tempfile::tempdir().unwrap();
7075 let mut config = McpConfig::default();
7076 let server = test_server_config();
7077 config.servers.insert("aws".to_string(), server.clone());
7078 let mut pool = McpPool::new(config);
7079 let sent = Arc::new(Mutex::new(Vec::new()));
7080 let mut connection = test_connection(Box::new(ScriptedValueTransport {
7081 sent: Arc::clone(&sent),
7082 responses: VecDeque::from([json_frame(serde_json::json!({
7083 "jsonrpc": "2.0",
7084 "id": 1,
7085 "error": { "code": -32000, "message": "UnauthorizedException (401): AWS token has expired" }
7086 }))]),
7087 }));
7088 connection.name = "aws".to_string();
7089 connection.config = server;
7090 connection.catalog_generation = pool.current_catalog_generation();
7091 connection.tools.push(McpTool {
7092 name: "lookup".to_string(),
7093 description: None,
7094 input_schema: serde_json::json!({"type": "object"}),
7095 annotations: None,
7096 });
7097 connection.resources.push(McpResource {
7098 uri: "aws://example".to_string(),
7099 name: "example".to_string(),
7100 description: None,
7101 mime_type: None,
7102 });
7103 connection.prompts.push(McpPrompt {
7104 name: "lookup".to_string(),
7105 description: None,
7106 arguments: Vec::new(),
7107 });
7108 pool.store_ready_connection("aws".to_string(), connection)
7109 .unwrap();
7110
7111 let result = match method {
7112 "tools/call" => {
7113 pool.call_tool("mcp_aws_lookup", serde_json::json!({}))
7114 .await
7115 }
7116 "resources/read" => pool.read_resource("aws", "aws://example").await,
7117 "prompts/get" => {
7118 pool.get_prompt("aws", "lookup", serde_json::json!({}))
7119 .await
7120 }
7121 _ => unreachable!(),
7122 };
7123 let error = result.expect_err("expired AWS credentials fail the call");
7124 assert!(format!("{error:#}").contains("aws sso login"));
7125 assert_eq!(sent.lock().unwrap().len(), 1, "never replay the call");
7126 assert!(!pool.connected_servers().contains(&"aws"));
7127 assert!(!pool.server_needs_auth("aws"));
7128 assert!(
7129 pool.to_api_tools()
7130 .iter()
7131 .all(|tool| !tool.name.starts_with("mcp_aws_"))
7132 );
7133
7134 // No external boot error map: the latest pool failure owns this row.
7135 let snapshot = pool.manager_snapshot(&dir.path().join("mcp.json"), false, &HashMap::new());
7136 let row = &snapshot.servers[0];
7137 assert!(!row.connected && !row.auth_required);
7138 assert_eq!(row.recovery_kind(false), Some(McpRecoveryKind::AwsLogin));
7139 assert!(
7140 row.error
7141 .as_deref()
7142 .unwrap_or_default()
7143 .contains("aws sso login")
7144 );
7145
7146 // Backoff reuses the recorded failure, so neither listing spawns a
7147 // process or logs in. Both model surfaces name the external recovery.
7148 for items in [
7149 pool.list_resources(None).await.unwrap(),
7150 pool.list_resource_templates(None).await.unwrap(),
7151 ] {
7152 let item = items
7153 .iter()
7154 .find(|item| item["server"] == "aws")
7155 .expect("AWS failure item");
7156 assert_eq!(item["error"], "aws_login_required");
7157 assert!(item.get("authenticate_tool").is_none());
7158 let message = item["message"].as_str().unwrap();
7159 assert!(message.contains("aws sso login") && message.contains("/mcp retry aws"));
7160 assert!(!message.contains("/mcp login"));
7161 }
7162 }
7163 }
7164
7165 #[test]
7166 fn sso_wording_on_an_oauth_capable_server_stays_on_the_oauth_route() {
7167 // Corporate SSO in front of an OAuth HTTP server: `aws sso login` cannot
7168 // help it, `/mcp login` can.
7169 assert!(!mcp_error_is_aws_login(
7170 "401 Unauthorized: SSO token expired",
7171 true
7172 ));
7173 assert_eq!(
7174 mcp_recovery_kind(
7175 true,
7176 true,
7177 false,
7178 Some("401 Unauthorized: SSO token expired"),
7179 true
7180 ),
7181 Some(McpRecoveryKind::Reauth)
7182 );
7183 assert_eq!(
7184 mcp_recovery_kind(true, true, false, Some("SSO token expired"), true),
7185 Some(McpRecoveryKind::Diagnose)
7186 );
7187 // The same text from a stdio server is still the AWS route.
7188 assert_eq!(
7189 mcp_recovery_kind(true, true, false, Some("SSO token expired"), false),
7190 Some(McpRecoveryKind::AwsLogin)
7191 );
7192
7193 let dir = tempfile::tempdir().unwrap();
7194 let mut config = McpConfig::default();
7195 let mut server = test_server_config();
7196 server.command = None;
7197 server.url = Some("https://mcp.example.test/mcp".to_string());
7198 // Automatic OAuth discovery works without explicit scopes/client config.
7199 assert!(mcp_server_oauth_capable(&server));
7200 config.servers.insert("corp".to_string(), server);
7201 let mut pool = McpPool::new(config);
7202 pool.note_connect_failure(
7203 "corp",
7204 &anyhow::anyhow!("401 Unauthorized: SSO token expired"),
7205 );
7206 assert!(pool.server_needs_auth("corp"));
7207 let errors = HashMap::new();
7208 let snapshot = pool.manager_snapshot(&dir.path().join("mcp.json"), false, &errors);
7209 let corp = &snapshot.servers[0];
7210 assert!(corp.auth_required);
7211 assert_eq!(corp.recovery_kind(true), Some(McpRecoveryKind::Reauth));
7212 }
7213
7214 /// Review of #6789: a reviewed plugin's initialize/live error is suppressed to
7215 /// protect environment-backed credentials, which also hid the AWS wording
7216 /// the recovery keys on. The classification now runs on the raw text first
7217 /// and only our fixed hint survives — never the server's own words, and
7218 /// never the plugin's `AWS_PROFILE` env value.
7219 #[tokio::test]
7220 async fn reviewed_plugin_rejection_keeps_aws_recovery_but_not_server_text() {
7221 let dir = tempfile::tempdir().unwrap();
7222 let plugins_root = dir.path().join("plugins");
7223 let plugin_base = plugins_root.join("aws-guard");
7224 fs::create_dir_all(&plugin_base).unwrap();
7225 fs::create_dir_all(dir.path().join("project")).unwrap();
7226 fs::write(
7227 plugin_base.join("plugin.toml"),
7228 "schema_version = 1\n[plugin]\nname = \"aws-guard\"\nversion = \"1.0.0\"\n",
7229 )
7230 .unwrap();
7231 let discovery = crate::plugins::discovery::DiscoveryConfig {
7232 workspace: dir.path().join("project"),
7233 user_plugins_dir: plugins_root,
7234 workspace_plugins_dir: dir.path().join("workspace-plugins-unused"),
7235 builtin_plugin_dirs: Vec::new(),
7236 state_path: dir.path().join("plugin-state/state.json"),
7237 };
7238 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
7239 registry.trust("aws-guard").unwrap();
7240 registry.enable("aws-guard").unwrap();
7241 let authority = registry.authority_for("aws-guard").unwrap();
7242
7243 let mut config = test_server_config();
7244 config
7245 .env
7246 .insert("AWS_PROFILE".to_string(), "env-profile-secret".to_string());
7247 config.reviewed_plugin = Some(
7248 ReviewedPluginMcpSource::from_authority(
7249 authority,
7250 None,
7251 Arc::new(crate::plugins::HostEnvironment::capture()),
7252 )
7253 .unwrap(),
7254 );
7255
7256 for initialize in [true, false] {
7257 for (server_error, login) in [
7258 (
7259 "Error retrieving credentials for acct-SECRET-123: The SSO session associated with this profile has expired or is otherwise invalid.",
7260 "aws sso login",
7261 ),
7262 (
7263 "acct-SECRET-123 LoginRefreshRequired: Please reauthenticate using aws login",
7264 "aws login",
7265 ),
7266 ] {
7267 let transport = ScriptedValueTransport {
7268 sent: Arc::new(Mutex::new(Vec::new())),
7269 responses: VecDeque::from([json_frame(serde_json::json!({
7270 "jsonrpc": "2.0",
7271 "id": 1,
7272 "error": { "code": -32602, "message": server_error }
7273 }))]),
7274 };
7275 let mut conn = test_connection(Box::new(transport));
7276 conn.config = config.clone();
7277 let error = if initialize {
7278 conn.initialize().await.expect_err("initialize rejected")
7279 } else {
7280 conn.call_method("tools/call", serde_json::json!({}), 5)
7281 .await
7282 .expect_err("live call rejected")
7283 };
7284 let error = format!("{error:#}");
7285 assert!(error.contains("server details suppressed"), "{error}");
7286 assert!(!error.contains("acct-SECRET-123"), "{error}");
7287 assert!(!error.contains("env-profile-secret"), "{error}");
7288 let hint = format!(
7289 "AWS credentials expired: run `{login}` in a terminal, then `/mcp retry mock`"
7290 );
7291 assert!(error.contains(&hint), "{error}");
7292 assert_eq!(
7293 mcp_recovery_kind(true, true, false, Some(&error), false),
7294 Some(McpRecoveryKind::AwsLogin)
7295 );
7296 // Pool/snapshot/listing classification sees only sanitized text;
7297 // it must preserve aws login versus aws sso login on that pass.
7298 assert_eq!(aws_login_hint(&config, "mock", &error), hint);
7299 }
7300 }
7301 }
7302
7303 #[test]
7304 fn mcp_recovery_kind_names_real_login_and_reload_commands() {
7305 assert_eq!(
7306 mcp_recovery_kind(false, true, false, None, false),
7307 Some(McpRecoveryKind::Enable)
7308 );
7309 assert_eq!(
7310 mcp_recovery_kind(true, false, false, None, false),
7311 Some(McpRecoveryKind::Connect)
7312 );
7313 assert_eq!(
7314 mcp_recovery_kind(true, true, false, Some("connection refused"), false),
7315 Some(McpRecoveryKind::Diagnose)
7316 );
7317 assert_eq!(
7318 mcp_recovery_kind(true, true, false, Some("connection refused"), true),
7319 Some(McpRecoveryKind::Diagnose)
7320 );
7321 assert_eq!(
7322 mcp_recovery_kind(true, true, false, Some("401 Unauthorized"), true),
7323 Some(McpRecoveryKind::Reauth)
7324 );
7325 assert_eq!(
7326 mcp_recovery_kind(true, true, false, None, true),
7327 Some(McpRecoveryKind::Reauth)
7328 );
7329 assert_eq!(
7330 mcp_recovery_kind(true, true, false, None, false),
7331 Some(McpRecoveryKind::Reconnect)
7332 );
7333 // Enabled, inspected, connected, no error: nothing to recover.
7334 assert_eq!(mcp_recovery_kind(true, true, true, None, false), None);
7335
7336 assert_eq!(
7337 McpRecoveryKind::Reauth.slash_command("github"),
7338 "/mcp login github"
7339 );
7340 // One row, one server. A `[reconnect] github` row that reloads every
7341 // configured server is not the action it named.
7342 assert_eq!(
7343 McpRecoveryKind::Connect.slash_command("github"),
7344 "/mcp retry github"
7345 );
7346 assert_eq!(
7347 McpRecoveryKind::Reconnect.slash_command("github"),
7348 "/mcp retry github"
7349 );
7350 // A name the command line cannot carry safely falls back to the blunt
7351 // reload rather than emitting an argument that would not survive parsing.
7352 assert_eq!(
7353 McpRecoveryKind::Reconnect.slash_command("name with spaces"),
7354 "/mcp reload"
7355 );
7356 assert_eq!(
7357 McpRecoveryKind::Diagnose.slash_command("github"),
7358 "/mcp validate github"
7359 );
7360 assert_eq!(
7361 McpRecoveryKind::Diagnose.slash_command("name with spaces"),
7362 "/mcp validate"
7363 );
7364 assert!(
7365 !McpRecoveryKind::Reauth
7366 .slash_command("github")
7367 .contains("/mcp auth")
7368 );
7369 assert!(mcp_name_is_command_safe("github"));
7370 assert!(!mcp_name_is_command_safe("github mcp"));
7371 }
7372
7373 // === Synthetic self-serve OAuth authenticate tool (agent self-serve auth) ===
7374 //
7375 // One loopback origin serving both the MCP endpoint and the OAuth
7376 // authorization-server APIs the login/refresh flows need:
7377 // `/.well-known/oauth-authorization-server*` metadata, RFC 7591 `/register`,
7378 // and a `/token` endpoint that rejects the `rt-stale` refresh token with
7379 // `invalid_grant` and accepts everything else. The MCP endpoint 401s until
7380 // the request carries `Authorization: Bearer cw-test-access`.
7381
7382 struct OAuthMcpMock {
7383 addr: std::net::SocketAddr,
7384 token_requests: Arc<AtomicUsize>,
7385 frames: Arc<Mutex<Vec<Value>>>,
7386 /// When set, the provider has revoked every grant: `/mcp` 401s even with
7387 /// the previously accepted bearer and `/token` rejects every refresh
7388 /// with `invalid_grant` — a mid-session revocation.
7389 revoked: Arc<std::sync::atomic::AtomicBool>,
7390 task: tokio::task::JoinHandle<()>,
7391 }
7392
7393 impl OAuthMcpMock {
7394 fn url(&self) -> String {
7395 format!("http://{}/mcp", self.addr)
7396 }
7397
7398 fn revoke_all_grants(&self) {
7399 self.revoked.store(true, AtomicOrdering::SeqCst);
7400 }
7401
7402 async fn spawn() -> Self {
7403 use tokio::io::{AsyncReadExt, AsyncWriteExt};
7404 use tokio::net::{TcpListener, TcpStream};
7405
7406 async fn read_request(socket: &mut TcpStream) -> String {
7407 let mut request = Vec::new();
7408 let mut buf = [0; 2048];
7409 let header_end = loop {
7410 let n = socket.read(&mut buf).await.unwrap();
7411 assert!(n > 0, "client closed before headers completed");
7412 request.extend_from_slice(&buf[..n]);
7413 if let Some(pos) = request.windows(4).position(|window| window == b"\r\n\r\n") {
7414 break pos + 4;
7415 }
7416 };
7417 let headers = String::from_utf8_lossy(&request[..header_end]);
7418 let content_length = headers
7419 .lines()
7420 .find_map(|line| {
7421 let (name, value) = line.split_once(':')?;
7422 name.eq_ignore_ascii_case("content-length")
7423 .then(|| value.trim().parse::<usize>().ok())
7424 .flatten()
7425 })
7426 .unwrap_or(0);
7427 let total_len = header_end + content_length;
7428 while request.len() < total_len {
7429 let n = socket.read(&mut buf).await.unwrap();
7430 assert!(n > 0, "client closed before body completed");
7431 request.extend_from_slice(&buf[..n]);
7432 }
7433 String::from_utf8(request).unwrap()
7434 }
7435
7436 async fn write_json(socket: &mut TcpStream, status: &str, body: &str) {
7437 let response = format!(
7438 "HTTP/1.1 {status}\r\nContent-Type: application/json\r\nConnection: close\r\nContent-Length: {}\r\n\r\n{body}",
7439 body.len()
7440 );
7441 socket.write_all(response.as_bytes()).await.unwrap();
7442 }
7443
7444 async fn write_empty(socket: &mut TcpStream, status: &str) {
7445 socket
7446 .write_all(
7447 format!("HTTP/1.1 {status}\r\nConnection: close\r\nContent-Length: 0\r\n\r\n")
7448 .as_bytes(),
7449 )
7450 .await
7451 .unwrap();
7452 }
7453
7454 async fn write_mcp_sse(
7455 socket: &mut TcpStream,
7456 id: serde_json::Value,
7457 result: serde_json::Value,
7458 ) {
7459 let payload = serde_json::json!({"jsonrpc": "2.0", "id": id, "result": result});
7460 let body = format!("event: message\ndata: {payload}\n\n");
7461 let response = format!(
7462 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\nContent-Length: {}\r\n\r\n{body}",
7463 body.len()
7464 );
7465 socket.write_all(response.as_bytes()).await.unwrap();
7466 }
7467
7468 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
7469 let addr = listener.local_addr().unwrap();
7470 let token_requests = Arc::new(AtomicUsize::new(0));
7471 let server_token_requests = Arc::clone(&token_requests);
7472 let revoked = Arc::new(std::sync::atomic::AtomicBool::new(false));
7473 let server_revoked = Arc::clone(&revoked);
7474 let frames = Arc::new(Mutex::new(Vec::new()));
7475 let server_frames = Arc::clone(&frames);
7476 let task = tokio::spawn(async move {
7477 loop {
7478 let Ok((mut socket, _)) = listener.accept().await else {
7479 break;
7480 };
7481 let token_requests = Arc::clone(&server_token_requests);
7482 let revoked = Arc::clone(&server_revoked);
7483 let frames = Arc::clone(&server_frames);
7484 tokio::spawn(async move {
7485 let request = read_request(&mut socket).await;
7486 let first_line = request.lines().next().unwrap_or("").to_string();
7487 let mut parts = first_line.split_whitespace();
7488 let method = parts.next().unwrap_or("").to_string();
7489 let path = parts.next().unwrap_or("").to_string();
7490 let path_only = path.split('?').next().unwrap_or("").to_string();
7491 let body = request.split("\r\n\r\n").nth(1).unwrap_or("").to_string();
7492 let revoked = revoked.load(AtomicOrdering::SeqCst);
7493 let authorized = !revoked
7494 && request
7495 .to_ascii_lowercase()
7496 .contains("authorization: bearer cw-test-access");
7497
7498 // RFC 8414: the authorization server lives at the origin
7499 // root, so it publishes its metadata only at the canonical
7500 // `/.well-known/oauth-authorization-server` and its
7501 // `issuer` is the root origin. The path-insertion
7502 // candidates rmcp probes first (`.../oauth-authorization-server/mcp`)
7503 // belong to a *different* issuer and must 404 here: rmcp
7504 // 3.2 validates the discovered `issuer` against the
7505 // discovery URL and rejects metadata served at the wrong
7506 // one.
7507 if method == "GET" && path_only == "/.well-known/oauth-authorization-server" {
7508 let metadata = serde_json::json!({
7509 "issuer": format!("http://{addr}"),
7510 "authorization_endpoint": format!("http://{addr}/authorize"),
7511 "token_endpoint": format!("http://{addr}/token"),
7512 "registration_endpoint": format!("http://{addr}/register"),
7513 "response_types_supported": ["code"],
7514 });
7515 write_json(&mut socket, "200 OK", &metadata.to_string()).await;
7516 } else if method == "GET" && path_only.starts_with("/.well-known/") {
7517 write_empty(&mut socket, "404 Not Found").await;
7518 } else if method == "POST" && path_only == "/register" {
7519 write_json(
7520 &mut socket,
7521 "200 OK",
7522 r#"{"client_id":"cw-test-client","redirect_uris":[]}"#,
7523 )
7524 .await;
7525 } else if method == "POST" && path_only == "/token" {
7526 token_requests.fetch_add(1, AtomicOrdering::SeqCst);
7527 if revoked || body.contains("refresh_token=rt-stale") {
7528 write_json(
7529 &mut socket,
7530 "400 Bad Request",
7531 r#"{"error":"invalid_grant","error_description":"stale grant"}"#,
7532 )
7533 .await;
7534 } else if body.contains("refresh_token=rt-rotated") {
7535 write_json(
7536 &mut socket,
7537 "200 OK",
7538 r#"{"access_token":"cw-rotated-access-2","token_type":"Bearer","expires_in":3600,"refresh_token":"rt-rotated-2"}"#,
7539 )
7540 .await;
7541 } else {
7542 write_json(
7543 &mut socket,
7544 "200 OK",
7545 r#"{"access_token":"cw-test-access","token_type":"Bearer","expires_in":3600,"refresh_token":"rt-fresh"}"#,
7546 )
7547 .await;
7548 }
7549 } else if path_only == "/mcp" {
7550 // The Streamable HTTP session preflight is a bodyless
7551 // GET; the real protocol starts at POST. 405 is the
7552 // spec-shaped refusal and never reaches the parser.
7553 if method == "GET" {
7554 write_empty(&mut socket, "405 Method Not Allowed").await;
7555 return;
7556 }
7557 let frame: Value = serde_json::from_str(&body).unwrap();
7558 frames.lock().unwrap().push(frame);
7559 if !authorized {
7560 write_empty(&mut socket, "401 Unauthorized").await;
7561 return;
7562 }
7563 let value: serde_json::Value = serde_json::from_str(&body).unwrap();
7564 let rpc_method = value["method"].as_str().unwrap_or("");
7565 if rpc_method == "notifications/initialized" {
7566 write_empty(&mut socket, "202 Accepted").await;
7567 return;
7568 }
7569 let id = value["id"].clone();
7570 let result = match rpc_method {
7571 "initialize" => serde_json::json!({
7572 "protocolVersion": "2024-11-05",
7573 "serverInfo": {"name": "mock-oauth", "version": "1.0.0"},
7574 "capabilities": {"tools": {}}
7575 }),
7576 "tools/list" => serde_json::json!({
7577 "tools": [{
7578 "name": "wiki_lookup",
7579 "description": "Look up a wiki page",
7580 "inputSchema": {"type": "object"}
7581 }]
7582 }),
7583 _ => serde_json::json!({}),
7584 };
7585 write_mcp_sse(&mut socket, id, result).await;
7586 } else {
7587 write_empty(&mut socket, "404 Not Found").await;
7588 }
7589 });
7590 }
7591 });
7592 OAuthMcpMock {
7593 addr,
7594 token_requests,
7595 frames,
7596 revoked,
7597 task,
7598 }
7599 }
7600 }
7601
7602 fn mock_oauth_server_config(addr: std::net::SocketAddr) -> McpServerConfig {
7603 let mut config = test_server_config();
7604 config.command = None;
7605 config.url = Some(format!("http://{addr}/mcp"));
7606 config
7607 }
7608
7609 fn seed_oauth_tokens(
7610 server_name: &str,
7611 url: &str,
7612 access_token: &str,
7613 refresh_token: &str,
7614 expires_at: Option<u64>,
7615 ) {
7616 let tokens: oauth::StoredMcpOAuthTokens = serde_json::from_value(serde_json::json!({
7617 "server_name": server_name,
7618 "url": url,
7619 "client_id": "cw-test-client",
7620 "token_response": {
7621 "access_token": access_token,
7622 "token_type": "Bearer",
7623 "refresh_token": refresh_token,
7624 "expires_in": 3600
7625 },
7626 "expires_at": expires_at
7627 }))
7628 .unwrap();
7629 oauth::save_oauth_tokens(&tokens).unwrap();
7630 }
7631
7632 fn millis_from_now(offset_ms: u64) -> u64 {
7633 (std::time::SystemTime::now() + Duration::from_millis(offset_ms))
7634 .duration_since(std::time::UNIX_EPOCH)
7635 .unwrap()
7636 .as_millis() as u64
7637 }
7638
7639 #[tokio::test]
7640 async fn model_reconnect_reuses_configured_credentials_without_restarting_siblings() {
7641 use crate::tools::runtime_mcp::StartRuntimeMcpServer;
7642 use crate::tools::spec::{ToolContext, ToolSpec};
7643
7644 let _env = crate::test_support::lock_test_env();
7645 let dir = tempfile::tempdir().unwrap();
7646 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
7647 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
7648 let _loopback = lock_mcp_loopback_tests().await;
7649 let mock = OAuthMcpMock::spawn().await;
7650 let name = "existing_server";
7651 let mut config = McpConfig::default();
7652 config
7653 .servers
7654 .insert(name.into(), mock_oauth_server_config(mock.addr));
7655 config
7656 .servers
7657 .insert("healthy".into(), mock_oauth_server_config(mock.addr));
7658 seed_oauth_tokens(
7659 "healthy",
7660 &mock.url(),
7661 "cw-test-access",
7662 "rt-fresh",
7663 Some(millis_from_now(3_600_000)),
7664 );
7665 let mut pool = McpPool::new(config);
7666 pool.get_or_connect("healthy").await.unwrap();
7667 let sibling_cancel = pool.connections["healthy"].cancel_token.clone();
7668 assert!(
7669 pool.get_or_connect(name).await.is_err(),
7670 "boot before login must require auth"
7671 );
7672 let pool = Arc::new(tokio::sync::Mutex::new(pool));
7673 let tool = StartRuntimeMcpServer::new(Arc::clone(&pool));
7674 let mut context = ToolContext::new(dir.path());
7675
7676 // Credentials arrive from the separate login process under the original key.
7677 seed_oauth_tokens(
7678 name,
7679 &mock.url(),
7680 "cw-test-access",
7681 "rt-fresh",
7682 Some(millis_from_now(3_600_000)),
7683 );
7684 context.disallowed_tools = vec![format!("mcp_{name}_*")];
7685 assert!(
7686 tool.execute(serde_json::json!({"name": name}), &context)
7687 .await
7688 .is_err()
7689 );
7690 assert!(!pool.lock().await.connected_servers().contains(&name));
7691 context.disallowed_tools.clear();
7692 let result = tool
7693 .execute(serde_json::json!({"name": name}), &context)
7694 .await
7695 .unwrap();
7696 assert_eq!(
7697 result.metadata,
7698 Some(serde_json::json!({"mcp_catalog_changed": true}))
7699 );
7700 let mut lock = pool.lock().await;
7701 assert!(lock.connected_servers().contains(&name));
7702 assert!(
7703 lock.all_tools()
7704 .iter()
7705 .any(|(tool, _)| tool == "mcp_existing_server_wiki_lookup")
7706 );
7707 assert!(
7708 lock.dynamic_servers.read().is_empty(),
7709 "reconnect cannot add an alias"
7710 );
7711 assert!(
7712 !sibling_cancel.is_cancelled(),
7713 "healthy sibling was restarted"
7714 );
7715 lock.call_tool("mcp_existing_server_wiki_lookup", serde_json::json!({}))
7716 .await
7717 .unwrap();
7718 drop(lock);
7719 assert!(
7720 tool.execute(serde_json::json!({"name": "absent"}), &context)
7721 .await
7722 .is_err()
7723 );
7724 assert!(pool.lock().await.dynamic_servers.read().is_empty());
7725 assert!(
7726 oauth::load_oauth_tokens("existing-server", &mock.url())
7727 .unwrap()
7728 .is_none(),
7729 "name must not be sanitized into another credential key"
7730 );
7731 mock.task.abort();
7732 }
7733
7734 #[tokio::test]
7735 async fn needs_auth_server_advertises_synthetic_authenticate_tool() {
7736 let _env = crate::test_support::lock_test_env();
7737 let dir = tempfile::tempdir().unwrap();
7738 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
7739 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
7740 let _loopback = lock_mcp_loopback_tests().await;
7741 let mock = OAuthMcpMock::spawn().await;
7742
7743 let mut mcp_config = McpConfig::default();
7744 mcp_config.servers.insert(
7745 "wikiserver".to_string(),
7746 mock_oauth_server_config(mock.addr),
7747 );
7748 // A server carrying manual bearer configuration that also 401s must not
7749 // get the OAuth tool: the flow is gated to servers OAuth can serve.
7750 let mut bearer_config = mock_oauth_server_config(mock.addr);
7751 bearer_config.bearer_token_env_var = Some("CW_TEST_UNSET_BEARER".to_string());
7752 mcp_config
7753 .servers
7754 .insert("beareronly".to_string(), bearer_config);
7755 let mut pool = McpPool::new(mcp_config);
7756
7757 let errors = pool.connect_all().await;
7758 assert_eq!(errors.len(), 2, "{errors:?}");
7759 for (name, err) in &errors {
7760 assert!(
7761 oauth::error_looks_auth_required(err),
7762 "{name} should fail auth-required: {err:#}"
7763 );
7764 }
7765 assert!(
7766 pool.all_tools().is_empty(),
7767 "needs-auth servers advertise no real tools"
7768 );
7769
7770 let tools = pool.to_api_tools();
7771 let auth_tool = tools
7772 .iter()
7773 .find(|tool| tool.name == "mcp_wikiserver_authenticate")
7774 .expect("needs-auth server must expose the synthetic authenticate tool");
7775 // The coaching contract: show the URL verbatim, the call blocks, real
7776 // tools replace the synthetic one on success.
7777 assert!(
7778 auth_tool
7779 .description
7780 .contains("shown to the user in the session status while this call waits"),
7781 "{}",
7782 auth_tool.description
7783 );
7784 assert!(
7785 auth_tool.description.contains("blocks (up to 5 minutes)"),
7786 "{}",
7787 auth_tool.description
7788 );
7789 assert!(
7790 auth_tool
7791 .description
7792 .contains("replace this synthetic authenticate tool"),
7793 "{}",
7794 auth_tool.description
7795 );
7796 assert!(
7797 auth_tool.description.contains("/mcp login wikiserver"),
7798 "{}",
7799 auth_tool.description
7800 );
7801 assert!(
7802 tools
7803 .iter()
7804 .all(|tool| !tool.name.starts_with("mcp_wikiserver_")
7805 || tool.name == "mcp_wikiserver_authenticate"),
7806 "no real wikiserver tools while needs-auth: {tools:?}"
7807 );
7808 assert!(
7809 tools
7810 .iter()
7811 .all(|tool| !tool.name.starts_with("mcp_beareronly_")),
7812 "a bearer-configured server must not get the OAuth tool: {tools:?}"
7813 );
7814
7815 // The execution predicate the engine consults: only the needs-auth
7816 // server's exact synthetic name resolves.
7817 assert_eq!(
7818 pool.authenticate_tool_target("mcp_wikiserver_authenticate")
7819 .as_deref(),
7820 Some("wikiserver")
7821 );
7822 assert_eq!(
7823 pool.authenticate_tool_target("mcp_beareronly_authenticate"),
7824 None,
7825 "manual bearer auth is not OAuth-servable"
7826 );
7827 assert_eq!(
7828 pool.authenticate_tool_target("mcp_wikiserver_wiki_lookup"),
7829 None
7830 );
7831 assert_eq!(
7832 pool.authenticate_tool_target("mcp_unknown_authenticate"),
7833 None
7834 );
7835
7836 // The typed `◆ auth required` state the TUI surfaces derive from the
7837 // same pool state: both 401 servers carry it, and the OAuth-servable one
7838 // routes to `/mcp login`.
7839 assert!(pool.server_needs_auth("wikiserver"));
7840 assert!(pool.server_needs_auth("beareronly"));
7841 let error_map: HashMap<String, String> = errors
7842 .iter()
7843 .map(|(name, err)| (name.clone(), format_mcp_error_for_display(err)))
7844 .collect();
7845 let snapshot = pool.manager_snapshot(&dir.path().join("mcp.json"), false, &error_map);
7846 let wiki = snapshot
7847 .servers
7848 .iter()
7849 .find(|server| server.name == "wikiserver")
7850 .expect("wikiserver in snapshot");
7851 assert!(wiki.auth_required, "{wiki:?}");
7852 assert!(!wiki.connected);
7853 let recovery = wiki
7854 .recovery_kind(false)
7855 .expect("a server needing auth has a recovery");
7856 assert_eq!(recovery, McpRecoveryKind::Reauth);
7857 assert_eq!(
7858 recovery.slash_command("wikiserver"),
7859 "/mcp login wikiserver"
7860 );
7861
7862 // A real tool name from a catalog built before the login lapsed must
7863 // not dead-end: the error names the synthetic tool that recovers it.
7864 let err = pool
7865 .call_tool("mcp_wikiserver_wiki_lookup", serde_json::json!({}))
7866 .await
7867 .expect_err("a needs-auth server cannot serve real tools");
7868 let text = format!("{err:#}");
7869 assert!(text.contains("mcp_wikiserver_authenticate"), "{text}");
7870 assert!(text.contains("◆ auth required"), "{text}");
7871 assert!(text.contains("/mcp login wikiserver"), "{text}");
7872 // Still needs-auth after the failed real call (no state churn).
7873 assert!(pool.server_needs_auth("wikiserver"));
7874
7875 mock.task.abort();
7876 }
7877
7878 #[tokio::test]
7879 async fn dead_refresh_grant_flips_server_to_auth_required_and_offers_login_tool() {
7880 use oauth2::TokenResponse as _;
7881
7882 let _env = crate::test_support::lock_test_env();
7883 let dir = tempfile::tempdir().unwrap();
7884 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
7885 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
7886 let _loopback = lock_mcp_loopback_tests().await;
7887
7888 let mock = OAuthMcpMock::spawn().await;
7889 let url = mock.url();
7890 let config = mock_oauth_server_config(mock.addr);
7891
7892 // A stored credential whose refresh grant the provider now rejects
7893 // (`invalid_grant`, nobody rotated it): the login has definitively
7894 // lapsed. Before this slice the server showed as a plain failure and
7895 // every reconnect replayed the same rejected refresh.
7896 seed_oauth_tokens("wikiserver", &url, "cw-stale-access", "rt-stale", Some(1));
7897 let stored = oauth::load_oauth_tokens("wikiserver", &url)
7898 .unwrap()
7899 .expect("seeded tokens");
7900 assert_eq!(
7901 stored.token_response.0.access_token().secret(),
7902 "cw-stale-access"
7903 );
7904
7905 let mut mcp_config = McpConfig::default();
7906 mcp_config.servers.insert("wikiserver".to_string(), config);
7907 let mut pool = McpPool::new(mcp_config);
7908 let errors = pool.connect_all().await;
7909 let (_, err) = errors
7910 .iter()
7911 .find(|(name, _)| name == "wikiserver")
7912 .expect("connect fails");
7913 assert!(
7914 oauth::error_looks_auth_required(err),
7915 "a dead grant classifies auth-required: {err:#}"
7916 );
7917 assert!(format!("{err:#}").contains("invalid_grant"), "{err:#}");
7918
7919 // The dead credential is invalidated so auth status no longer claims
7920 // "logged in", the typed state flips, and the model gets the login tool.
7921 assert!(
7922 oauth::load_oauth_tokens("wikiserver", &url)
7923 .unwrap()
7924 .is_none(),
7925 "a definitively rejected grant is removed from the store"
7926 );
7927 assert!(pool.server_needs_auth("wikiserver"));
7928 assert_eq!(
7929 pool.authenticate_tool_target("mcp_wikiserver_authenticate")
7930 .as_deref(),
7931 Some("wikiserver")
7932 );
7933 assert!(
7934 pool.to_api_tools()
7935 .iter()
7936 .any(|tool| tool.name == "mcp_wikiserver_authenticate"),
7937 "dead grant must offer the self-serve login tool"
7938 );
7939 assert_eq!(
7940 mock.token_requests.load(AtomicOrdering::SeqCst),
7941 1,
7942 "one rejected refresh; no retry against an unchanged store"
7943 );
7944
7945 mock.task.abort();
7946 }
7947
7948 #[tokio::test]
7949 async fn selfserve_auth_flow_persists_tokens_and_swaps_real_tools_back() {
7950 use oauth2::TokenResponse as _;
7951
7952 let _env = crate::test_support::lock_test_env();
7953 let dir = tempfile::tempdir().unwrap();
7954 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
7955 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
7956 let _loopback = lock_mcp_loopback_tests().await;
7957
7958 let mock = OAuthMcpMock::spawn().await;
7959 let url = mock.url();
7960 let config = mock_oauth_server_config(mock.addr);
7961
7962 // Before any login, the pool's connect fails auth-required and the model
7963 // catalog exposes the synthetic tool in place of the unavailable real
7964 // tools.
7965 let mut mcp_config = McpConfig::default();
7966 mcp_config
7967 .servers
7968 .insert("wikiserver".to_string(), config.clone());
7969 let mut pool = McpPool::new(mcp_config);
7970 let errors = pool.connect_all().await;
7971 assert!(
7972 errors
7973 .iter()
7974 .any(|(name, err)| name == "wikiserver" && oauth::error_looks_auth_required(err)),
7975 "{errors:?}"
7976 );
7977 assert!(pool.all_tools().is_empty());
7978 assert!(
7979 pool.to_api_tools()
7980 .iter()
7981 .any(|tool| tool.name == "mcp_wikiserver_authenticate")
7982 );
7983
7984 // A declined browser flow must surface a truthful error the model can
7985 // relay, and leave the server in needs-auth.
7986 let login =
7987 oauth::begin_oauth_login_for_server_tool("wikiserver", &config, None, None, None, None)
7988 .await
7989 .unwrap();
7990 let auth_url = reqwest::Url::parse(login.authorization_url()).unwrap();
7991 let redirect_uri = auth_url
7992 .query_pairs()
7993 .find(|(key, _)| key == "redirect_uri")
7994 .map(|(_, value)| value.into_owned())
7995 .expect("authorization URL carries redirect_uri");
7996 let flow = tokio::spawn(login.finish());
7997 test_http_client()
7998 .get(format!(
7999 "{redirect_uri}?error=access_denied&error_description=not-today"
8000 ))
8001 .send()
8002 .await
8003 .unwrap();
8004 let err = flow
8005 .await
8006 .unwrap()
8007 .expect_err("a declined flow must error, not silently succeed");
8008 assert!(format!("{err:#}").contains("access_denied"), "{err:#}");
8009 assert!(
8010 oauth::load_oauth_tokens("wikiserver", &url)
8011 .unwrap()
8012 .is_none(),
8013 "a declined flow persists no tokens"
8014 );
8015
8016 // The approved flow: drive the loopback callback in-test.
8017 let login =
8018 oauth::begin_oauth_login_for_server_tool("wikiserver", &config, None, None, None, None)
8019 .await
8020 .unwrap();
8021 let auth_url = reqwest::Url::parse(login.authorization_url()).unwrap();
8022 let state = auth_url
8023 .query_pairs()
8024 .find(|(key, _)| key == "state")
8025 .map(|(_, value)| value.into_owned())
8026 .expect("authorization URL carries state");
8027 let redirect_uri = auth_url
8028 .query_pairs()
8029 .find(|(key, _)| key == "redirect_uri")
8030 .map(|(_, value)| value.into_owned())
8031 .expect("authorization URL carries redirect_uri");
8032 let flow = tokio::spawn(login.finish());
8033 test_http_client()
8034 .get(format!("{redirect_uri}?code=cw-test-code&state={state}"))
8035 .send()
8036 .await
8037 .unwrap();
8038 flow.await.unwrap().expect("approved flow completes");
8039
8040 let stored = oauth::load_oauth_tokens("wikiserver", &url)
8041 .unwrap()
8042 .expect("successful flow persists tokens");
8043 assert_eq!(
8044 stored.token_response.0.access_token().secret(),
8045 "cw-test-access"
8046 );
8047
8048 // The synthetic tool now adopts the completed login and swaps the real
8049 // tools back through call_tool (the already-authorized branch, since the
8050 // flow above persisted tokens to the shared store).
8051 let result = pool
8052 .call_tool("mcp_wikiserver_authenticate", serde_json::json!({}))
8053 .await
8054 .unwrap();
8055 assert_eq!(result["status"], "already_authorized", "{result}");
8056 assert!(
8057 result["tools"]
8058 .as_array()
8059 .unwrap()
8060 .contains(&serde_json::json!("mcp_wikiserver_wiki_lookup")),
8061 "{result}"
8062 );
8063
8064 let real_names: Vec<String> = pool
8065 .all_tools()
8066 .iter()
8067 .map(|(name, _)| name.clone())
8068 .collect();
8069 assert_eq!(real_names, vec!["mcp_wikiserver_wiki_lookup".to_string()]);
8070 let catalog = pool.to_api_tools();
8071 assert!(
8072 catalog
8073 .iter()
8074 .any(|tool| tool.name == "mcp_wikiserver_wiki_lookup"),
8075 "real tools must be in the catalog after auth: {catalog:?}"
8076 );
8077 assert!(
8078 catalog
8079 .iter()
8080 .all(|tool| tool.name != "mcp_wikiserver_authenticate"),
8081 "a connected server no longer advertises the synthetic tool: {catalog:?}"
8082 );
8083
8084 mock.task.abort();
8085 }
8086
8087 #[tokio::test]
8088 async fn invalid_grant_refresh_adopts_rotated_on_disk_token() {
8089 let _env = crate::test_support::lock_test_env();
8090 let dir = tempfile::tempdir().unwrap();
8091 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
8092 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
8093 let _loopback = lock_mcp_loopback_tests().await;
8094
8095 let mock = OAuthMcpMock::spawn().await;
8096 let url = mock.url();
8097 let config = mock_oauth_server_config(mock.addr);
8098
8099 // This runtime loaded a stale credential; its refresh grant fails.
8100 seed_oauth_tokens("wikiserver", &url, "cw-stale-access", "rt-stale", Some(1));
8101 let runtime = oauth::McpOAuthRuntime::from_server_config(
8102 "wikiserver",
8103 &config,
8104 reqwest::header::HeaderMap::new(),
8105 )
8106 .await
8107 .unwrap()
8108 .expect("stored tokens produce a runtime");
8109 // Another process rotates the on-disk credential before our refresh runs.
8110 seed_oauth_tokens(
8111 "wikiserver",
8112 &url,
8113 "cw-rotated-access",
8114 "rt-rotated",
8115 Some(millis_from_now(3_600_000)),
8116 );
8117
8118 let header = runtime
8119 .authorization_header()
8120 .await
8121 .unwrap()
8122 .expect("adopted token produces a header");
8123 assert_eq!(header, "Bearer cw-rotated-access");
8124 assert_eq!(
8125 mock.token_requests.load(AtomicOrdering::SeqCst),
8126 1,
8127 "a fresh rotated token is adopted without burning a second refresh"
8128 );
8129
8130 mock.task.abort();
8131 }
8132
8133 #[tokio::test]
8134 async fn invalid_grant_refresh_retries_once_with_rotated_grant() {
8135 use oauth2::TokenResponse as _;
8136
8137 let _env = crate::test_support::lock_test_env();
8138 let dir = tempfile::tempdir().unwrap();
8139 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
8140 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
8141 let _loopback = lock_mcp_loopback_tests().await;
8142
8143 let mock = OAuthMcpMock::spawn().await;
8144 let url = mock.url();
8145 let config = mock_oauth_server_config(mock.addr);
8146
8147 seed_oauth_tokens("wikiserver", &url, "cw-stale-access", "rt-stale", Some(1));
8148 let runtime = oauth::McpOAuthRuntime::from_server_config(
8149 "wikiserver",
8150 &config,
8151 reqwest::header::HeaderMap::new(),
8152 )
8153 .await
8154 .unwrap()
8155 .expect("stored tokens produce a runtime");
8156 // The rotated credential is itself expired, so the refresh must be
8157 // retried once with the rotated grant.
8158 seed_oauth_tokens(
8159 "wikiserver",
8160 &url,
8161 "cw-rotated-stale",
8162 "rt-rotated",
8163 Some(1),
8164 );
8165
8166 let header = runtime
8167 .authorization_header()
8168 .await
8169 .unwrap()
8170 .expect("retry with the rotated grant succeeds");
8171 assert_eq!(header, "Bearer cw-rotated-access-2");
8172 assert_eq!(
8173 mock.token_requests.load(AtomicOrdering::SeqCst),
8174 2,
8175 "stale grant failed once, rotated grant retried once"
8176 );
8177 let stored = oauth::load_oauth_tokens("wikiserver", &url)
8178 .unwrap()
8179 .expect("refreshed tokens persisted");
8180 assert_eq!(
8181 stored.token_response.0.access_token().secret(),
8182 "cw-rotated-access-2"
8183 );
8184
8185 mock.task.abort();
8186 }
8187
8188 #[tokio::test]
8189 async fn invalid_grant_refresh_with_unchanged_store_invalidates_dead_grant() {
8190 let _env = crate::test_support::lock_test_env();
8191 let dir = tempfile::tempdir().unwrap();
8192 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
8193 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
8194 let _loopback = lock_mcp_loopback_tests().await;
8195
8196 let mock = OAuthMcpMock::spawn().await;
8197 let url = mock.url();
8198 let config = mock_oauth_server_config(mock.addr);
8199
8200 seed_oauth_tokens("wikiserver", &url, "cw-stale-access", "rt-stale", Some(1));
8201 let runtime = oauth::McpOAuthRuntime::from_server_config(
8202 "wikiserver",
8203 &config,
8204 reqwest::header::HeaderMap::new(),
8205 )
8206 .await
8207 .unwrap()
8208 .expect("stored tokens produce a runtime");
8209
8210 let err = runtime
8211 .authorization_header()
8212 .await
8213 .expect_err("an unchanged store must surface the provider error");
8214 assert!(format!("{err:#}").contains("invalid_grant"), "{err:#}");
8215 assert!(
8216 oauth::error_looks_auth_required(&err),
8217 "a dead grant is an auth-required failure: {err:#}"
8218 );
8219 assert!(
8220 oauth::load_oauth_tokens("wikiserver", &url)
8221 .unwrap()
8222 .is_none(),
8223 "a grant the provider definitively rejected, that nobody rotated, is invalidated"
8224 );
8225 assert_eq!(
8226 mock.token_requests.load(AtomicOrdering::SeqCst),
8227 1,
8228 "no retry when the on-disk credential is unchanged"
8229 );
8230
8231 mock.task.abort();
8232 }
8233
8234 #[tokio::test]
8235 async fn invalidation_never_deletes_a_credential_rotated_after_the_re_read() {
8236 use oauth2::TokenResponse as _;
8237
8238 let _env = crate::test_support::lock_test_env();
8239 let dir = tempfile::tempdir().unwrap();
8240 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
8241 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
8242 let _loopback = lock_mcp_loopback_tests().await;
8243
8244 let mock = OAuthMcpMock::spawn().await;
8245 let url = mock.url();
8246 let config = mock_oauth_server_config(mock.addr);
8247
8248 // Both the held and the rotated credential are dead grants, so the
8249 // rotated one is adopted, retried, and rejected too. The runtime then
8250 // invalidates — but a peer that wrote a third credential in between
8251 // owns the durable winner, which must survive.
8252 seed_oauth_tokens("wikiserver", &url, "cw-stale-access", "rt-stale", Some(1));
8253 let runtime = oauth::McpOAuthRuntime::from_server_config(
8254 "wikiserver",
8255 &config,
8256 reqwest::header::HeaderMap::new(),
8257 )
8258 .await
8259 .unwrap()
8260 .expect("stored tokens produce a runtime");
8261 let err = runtime
8262 .authorization_header()
8263 .await
8264 .expect_err("dead grant fails");
8265 assert!(format!("{err:#}").contains("invalid_grant"), "{err:#}");
8266 assert!(
8267 oauth::load_oauth_tokens("wikiserver", &url)
8268 .unwrap()
8269 .is_none(),
8270 "unchanged store: invalidated"
8271 );
8272
8273 // Now the peer-wins case: a different credential lands on disk before
8274 // the (second) runtime invalidates its own stale copy.
8275 seed_oauth_tokens("wikiserver", &url, "cw-stale-access", "rt-stale", Some(1));
8276 let runtime = oauth::McpOAuthRuntime::from_server_config(
8277 "wikiserver",
8278 &config,
8279 reqwest::header::HeaderMap::new(),
8280 )
8281 .await
8282 .unwrap()
8283 .expect("stored tokens produce a runtime");
8284 // Peer rotates to a fresh, non-expired credential the mock accepts.
8285 seed_oauth_tokens(
8286 "wikiserver",
8287 &url,
8288 "cw-peer-access",
8289 "rt-fresh",
8290 Some(millis_from_now(3_600_000)),
8291 );
8292 let header = runtime
8293 .authorization_header()
8294 .await
8295 .unwrap()
8296 .expect("adopts the peer's fresh credential");
8297 assert_eq!(header, "Bearer cw-peer-access");
8298 let stored = oauth::load_oauth_tokens("wikiserver", &url)
8299 .unwrap()
8300 .expect("the peer's credential is still on disk");
8301 assert_eq!(
8302 stored
8303 .token_response
8304 .0
8305 .refresh_token()
8306 .map(|t| t.secret().as_str()),
8307 Some("rt-fresh")
8308 );
8309
8310 mock.task.abort();
8311 }
8312
8313 #[tokio::test]
8314 async fn mid_session_revocation_lands_in_the_same_auth_required_state() {
8315 mid_session_revocation_for_backend(McpBackend::Rust).await;
8316 }
8317 #[tokio::test(flavor = "current_thread")]
8318 async fn host_mid_session_revocation_lands_in_the_same_auth_required_state_without_replay() {
8319 mid_session_revocation_for_backend(McpBackend::Host).await;
8320 }
8321 async fn mid_session_revocation_for_backend(backend: McpBackend) {
8322 let _env = crate::test_support::lock_test_env();
8323 let dir = tempfile::tempdir().unwrap();
8324 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
8325 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
8326 let _loopback = lock_mcp_loopback_tests().await;
8327 let _policy = crate::plugins::activation::TestPolicyGuard::extension_host(true);
8328 let manager = (backend == McpBackend::Host).then(|| {
8329 let node = crate::extension_host::tests::node_for_tests("Host revoked OAuth parity")
8330 .expect("Host parity requires Node");
8331 Arc::new(crate::extension_host::ExtensionHostManager::new(
8332 crate::extension_host::ExtensionHostOptions {
8333 root: Some(dir.path().to_path_buf()),
8334 node_override: Some(node),
8335 ..Default::default()
8336 },
8337 ))
8338 });
8339 let _manager = manager
8340 .as_ref()
8341 .map(|manager| crate::extension_host::TestManagerGuard::install(Arc::clone(manager)));
8342
8343 let mock = OAuthMcpMock::spawn().await;
8344 let url = mock.url();
8345 let config = mock_oauth_server_config(mock.addr);
8346
8347 // A healthy session: stored credential accepted, real tools advertised.
8348 seed_oauth_tokens(
8349 "wikiserver",
8350 &url,
8351 "cw-test-access",
8352 "rt-fresh",
8353 Some(millis_from_now(3_600_000)),
8354 );
8355 let mut mcp_config = McpConfig::default();
8356 mcp_config.servers.insert("wikiserver".to_string(), config);
8357 let mut pool = McpPool::new(mcp_config).with_backend(backend);
8358 let errors = pool.connect_all().await;
8359 assert!(errors.is_empty(), "{errors:?}");
8360 assert!(!pool.server_needs_auth("wikiserver"));
8361 assert!(
8362 pool.to_api_tools()
8363 .iter()
8364 .any(|tool| tool.name == "mcp_wikiserver_wiki_lookup")
8365 );
8366
8367 // The provider revokes every grant mid-session: the live call 401s, the
8368 // reactive refresh is rejected with invalid_grant, and the failure must
8369 // land in the same typed state a failed connect produces — not a dead
8370 // transport error on a connection the pool still calls "ready".
8371 // Cross a whole second: reloading the same durable credential now has a
8372 // smaller derived expires_in, which must not look like a peer rotation.
8373 tokio::time::sleep(Duration::from_millis(1_100)).await;
8374 mock.revoke_all_grants();
8375 let err = pool
8376 .call_tool("mcp_wikiserver_wiki_lookup", serde_json::json!({}))
8377 .await
8378 .expect_err("a revoked credential cannot serve real tools");
8379 let text = format!("{err:#}");
8380 assert!(text.contains("◆ auth required"), "{text}");
8381 assert!(text.contains("mcp_wikiserver_authenticate"), "{text}");
8382 assert!(pool.server_needs_auth("wikiserver"));
8383 assert!(
8384 oauth::load_oauth_tokens("wikiserver", &url)
8385 .unwrap()
8386 .is_none(),
8387 "the definitively rejected credential is invalidated; refresh requests: {}; failure: {text}",
8388 mock.token_requests
8389 .load(std::sync::atomic::Ordering::SeqCst)
8390 );
8391 let catalog = pool.to_api_tools();
8392 assert!(
8393 catalog
8394 .iter()
8395 .any(|tool| tool.name == "mcp_wikiserver_authenticate"),
8396 "{catalog:?}"
8397 );
8398 assert!(
8399 catalog
8400 .iter()
8401 .all(|tool| tool.name != "mcp_wikiserver_wiki_lookup"),
8402 "a dropped connection advertises no real tools: {catalog:?}"
8403 );
8404
8405 // Listing surfaces carry the same recovery, naming the synthetic tool.
8406 let resources = pool.list_resources(None).await.unwrap();
8407 let item = resources
8408 .iter()
8409 .find(|item| item["error"] == "authentication_required")
8410 .expect("needs-auth server yields an auth-required listing item");
8411 assert_eq!(item["server"], "wikiserver");
8412 assert_eq!(item["authenticate_tool"], "mcp_wikiserver_authenticate");
8413 assert!(
8414 item["message"]
8415 .as_str()
8416 .unwrap()
8417 .contains("mcp_wikiserver_authenticate"),
8418 "{item}"
8419 );
8420
8421 // Every TUI surface derives from the same typed state.
8422 let snapshot = pool.manager_snapshot(&dir.path().join("mcp.json"), false, &HashMap::new());
8423 let wiki = snapshot
8424 .servers
8425 .iter()
8426 .find(|server| server.name == "wikiserver")
8427 .expect("wikiserver in snapshot");
8428 assert!(wiki.auth_required, "{wiki:?}");
8429 assert!(!wiki.connected);
8430 assert_eq!(wiki.recovery_kind(false), Some(McpRecoveryKind::Reauth));
8431
8432 assert_eq!(
8433 mock.frames
8434 .lock()
8435 .unwrap()
8436 .iter()
8437 .filter(|frame| frame["method"] == "tools/call")
8438 .count(),
8439 1,
8440 "rejected refresh must not replay a tool"
8441 );
8442 pool.shutdown_all().await;
8443 if let Some(manager) = manager {
8444 manager.shutdown().await;
8445 }
8446 mock.task.abort();
8447 }
8448
8449 #[tokio::test]
8450 async fn authenticate_tool_via_pool_releases_the_lock_during_the_browser_wait() {
8451 let _env = crate::test_support::lock_test_env();
8452 let dir = tempfile::tempdir().unwrap();
8453 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
8454 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
8455 let _loopback = lock_mcp_loopback_tests().await;
8456
8457 let mock = OAuthMcpMock::spawn().await;
8458 let mut mcp_config = McpConfig::default();
8459 mcp_config.servers.insert(
8460 "wikiserver".to_string(),
8461 mock_oauth_server_config(mock.addr),
8462 );
8463 let pool = Arc::new(tokio::sync::Mutex::new(McpPool::new(mcp_config)));
8464 let errors = pool.lock().await.connect_all().await;
8465 assert!(
8466 errors
8467 .iter()
8468 .any(|(name, err)| name == "wikiserver" && oauth::error_looks_auth_required(err)),
8469 "{errors:?}"
8470 );
8471
8472 // The engine path: the URL reaches the runtime before the wait, and the
8473 // pool lock is free while the user signs in.
8474 let (url_tx, url_rx) = tokio::sync::oneshot::channel();
8475 let flow = tokio::spawn({
8476 let pool = Arc::clone(&pool);
8477 async move {
8478 authenticate_tool_via_pool(&pool, "wikiserver", |url| {
8479 let _ = url_tx.send(url.to_string());
8480 })
8481 .await
8482 }
8483 });
8484 let auth_url = url_rx
8485 .await
8486 .expect("authorization URL announced before the wait");
8487 {
8488 let guard = tokio::time::timeout(Duration::from_secs(5), pool.lock())
8489 .await
8490 .expect("pool lock must not be held during the browser wait");
8491 assert!(guard.server_needs_auth("wikiserver"));
8492 assert!(
8493 guard
8494 .to_api_tools()
8495 .iter()
8496 .any(|tool| tool.name == "mcp_wikiserver_authenticate"),
8497 "still needs-auth until the flow completes"
8498 );
8499 }
8500
8501 let parsed = reqwest::Url::parse(&auth_url).unwrap();
8502 let state = parsed
8503 .query_pairs()
8504 .find(|(key, _)| key == "state")
8505 .map(|(_, value)| value.into_owned())
8506 .expect("authorization URL carries state");
8507 let redirect_uri = parsed
8508 .query_pairs()
8509 .find(|(key, _)| key == "redirect_uri")
8510 .map(|(_, value)| value.into_owned())
8511 .expect("authorization URL carries redirect_uri");
8512 test_http_client()
8513 .get(format!("{redirect_uri}?code=cw-test-code&state={state}"))
8514 .send()
8515 .await
8516 .unwrap();
8517 let result = flow
8518 .await
8519 .unwrap()
8520 .expect("approved flow reconnects the server");
8521 assert_eq!(result["status"], "authenticated", "{result}");
8522 assert_eq!(result["authorization_url"], auth_url, "{result}");
8523 assert!(
8524 result["tools"]
8525 .as_array()
8526 .unwrap()
8527 .contains(&serde_json::json!("mcp_wikiserver_wiki_lookup")),
8528 "{result}"
8529 );
8530
8531 let guard = pool.lock().await;
8532 assert!(!guard.server_needs_auth("wikiserver"));
8533 let catalog = guard.to_api_tools();
8534 assert!(
8535 catalog
8536 .iter()
8537 .any(|tool| tool.name == "mcp_wikiserver_wiki_lookup"),
8538 "{catalog:?}"
8539 );
8540 assert!(
8541 catalog
8542 .iter()
8543 .all(|tool| tool.name != "mcp_wikiserver_authenticate"),
8544 "{catalog:?}"
8545 );
8546 drop(guard);
8547
8548 mock.task.abort();
8549 }
8550
8551 #[tokio::test]
8552 async fn plugin_contributed_server_auth_required_names_its_env_credential_not_oauth_login() {
8553 let _env = crate::test_support::lock_test_env();
8554 let dir = tempfile::tempdir().unwrap();
8555 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
8556 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
8557 let _loopback = lock_mcp_loopback_tests().await;
8558
8559 let plugin_base = dir.path().join("plugins/wiki-plugin");
8560 fs::create_dir_all(&plugin_base).unwrap();
8561 fs::write(
8562 plugin_base.join("plugin.toml"),
8563 "schema_version = 1\n[plugin]\nname = \"wiki-plugin\"\nversion = \"1.0.0\"\n",
8564 )
8565 .unwrap();
8566 let (_, authority) = active_plugin_fixture(&plugin_base);
8567
8568 // A plugin-contributed server whose reviewed environment-backed header is
8569 // absent, against an endpoint that 401s without it.
8570 let mock = OAuthMcpMock::spawn().await;
8571 let endpoint = mock.url();
8572 let mut server = mock_oauth_server_config(mock.addr);
8573 server.env_headers.insert(
8574 "Authorization".to_string(),
8575 "CW_TEST_WIKI_PLUGIN_TOKEN".to_string(),
8576 );
8577 server.reviewed_plugin = Some(
8578 ReviewedPluginMcpSource::from_authority(
8579 authority,
8580 Some(&endpoint),
8581 Arc::new(crate::plugins::HostEnvironment::default()),
8582 )
8583 .unwrap(),
8584 );
8585 let mut mcp_config = McpConfig::default();
8586 mcp_config.servers.insert("wikiplugin".to_string(), server);
8587 let mut pool = McpPool::new(mcp_config);
8588
8589 let errors = pool.connect_all().await;
8590 let (_, err) = errors
8591 .iter()
8592 .find(|(name, _)| name == "wikiplugin")
8593 .expect("plugin server fails to connect");
8594 assert!(oauth::error_looks_auth_required(err), "{err:#}");
8595
8596 // Same typed state as an OAuth server, but never the OAuth tool: plugin
8597 // servers authenticate through their reviewed environment header.
8598 assert!(pool.server_needs_auth("wikiplugin"));
8599 assert_eq!(
8600 pool.authenticate_tool_target("mcp_wikiplugin_authenticate"),
8601 None
8602 );
8603 assert!(
8604 pool.to_api_tools()
8605 .iter()
8606 .all(|tool| !tool.name.starts_with("mcp_wikiplugin_")),
8607 "no synthetic OAuth tool for a plugin-contributed server"
8608 );
8609
8610 let err = pool
8611 .call_tool("mcp_wikiplugin_wiki_lookup", serde_json::json!({}))
8612 .await
8613 .expect_err("needs-auth plugin server cannot serve real tools");
8614 let text = format!("{err:#}");
8615 assert!(text.contains("◆ auth required"), "{text}");
8616 assert!(text.contains("plugin 'wiki-plugin'"), "{text}");
8617 assert!(text.contains("CW_TEST_WIKI_PLUGIN_TOKEN"), "{text}");
8618 assert!(text.contains("/mcp reload"), "{text}");
8619 assert!(!text.contains("/mcp login"), "{text}");
8620 assert!(!text.contains("mcp_wikiplugin_authenticate"), "{text}");
8621
8622 let snapshot = pool.manager_snapshot(&dir.path().join("mcp.json"), false, &HashMap::new());
8623 let plugin = snapshot
8624 .servers
8625 .iter()
8626 .find(|server| server.name == "wikiplugin")
8627 .expect("plugin server in snapshot");
8628 assert!(plugin.auth_required, "{plugin:?}");
8629
8630 mock.task.abort();
8631 }
8632
8633 #[test]
8634 fn mcp_display_target_shows_command_names_only() {
8635 // stdio: no `./…` path prefix, no directories, no args.
8636 assert_eq!(
8637 mcp_display_target("stdio", "./mcp/custom-server --port 8080"),
8638 "custom-server"
8639 );
8640 assert_eq!(mcp_display_target("stdio", "node server.js"), "node");
8641 assert_eq!(mcp_display_target("stdio", "/usr/local/bin/foo -x"), "foo");
8642 assert_eq!(
8643 mcp_display_target("stdio", "C:\\tools\\mcp.exe --stdio"),
8644 "mcp.exe"
8645 );
8646 assert_eq!(mcp_display_target("stdio", "(missing)"), "(missing)");
8647 // URL transports keep the full URL: it is the identity.
8648 assert_eq!(
8649 mcp_display_target("http/sse", "https://example.invalid/mcp"),
8650 "https://example.invalid/mcp"
8651 );
8652 assert_eq!(
8653 mcp_display_target("sse", "https://example.invalid/sse?token=abc"),
8654 "https://example.invalid/sse?token=abc"
8655 );
8656 }
8657
8658 fn test_mcp_http_client(url: &str) -> super::http_client::McpHttpClient {
8659 super::http_client::McpHttpClient::new(
8660 url,
8661 false,
8662 false,
8663 false,
8664 None,
8665 Duration::from_secs(10),
8666 Duration::from_secs(120),
8667 )
8668 .expect("MCP fixture client")
8669 }
8670
8671 fn ceiling_test_connection(name: &str, sent: Arc<Mutex<Vec<serde_json::Value>>>) -> McpConnection {
8672 let mut connection = test_connection(Box::new(ScriptedValueTransport {
8673 sent,
8674 responses: VecDeque::from([json_frame(serde_json::json!({
8675 "jsonrpc": "2.0", "id": 1, "result": {"ok": true}
8676 }))]),
8677 }));
8678 connection.name = name.to_string();
8679 connection.tools = ["read", "delete"]
8680 .into_iter()
8681 .map(|name| McpTool {
8682 name: name.to_string(),
8683 description: None,
8684 input_schema: serde_json::json!({}),
8685 annotations: None,
8686 })
8687 .collect();
8688 connection.resources = vec![McpResource {
8689 name: "one".to_string(),
8690 uri: "memory://one".to_string(),
8691 description: None,
8692 mime_type: None,
8693 }];
8694 connection.resource_templates = vec![McpResourceTemplate {
8695 name: "items".to_string(),
8696 uri_template: "memory://{id}".to_string(),
8697 description: None,
8698 mime_type: None,
8699 }];
8700 connection.prompts = vec![McpPrompt {
8701 name: "review".to_string(),
8702 description: None,
8703 arguments: vec![],
8704 }];
8705 connection
8706 }
8707
8708 #[tokio::test]
8709 async fn mcp_ceiling_denied_server_is_absent_across_cached_boot_meta_auth_and_runtime_paths() {
8710 let sent = Arc::new(Mutex::new(Vec::new()));
8711 let connection = ceiling_test_connection("private_a", Arc::clone(&sent));
8712 let mut config = connection.config.clone();
8713 config.required = true;
8714 config.url = Some("https://mcp.example.com".to_string());
8715 config.command = None;
8716 config.scopes = vec!["read".to_string()];
8717 let mut pool = McpPool::new(McpConfig {
8718 servers: HashMap::from([("private_a".to_string(), config.clone())]),
8719 timeouts: McpTimeouts::default(),
8720 })
8721 .with_disallowed_tools(vec!["MCP_PRIVATE_A_*".to_string()]);
8722 // A previously connected or auth-failed entry must not become reachable.
8723 pool.connections.insert("private_a".to_string(), connection);
8724 pool.needs_auth_servers.insert("private_a".to_string());
8725 assert!(pool.all_tools().is_empty());
8726 assert!(pool.all_resources().is_empty());
8727 assert!(pool.all_resource_templates().is_empty());
8728 assert!(pool.all_prompts().is_empty());
8729 assert!(pool.resolved_tool_servers().is_empty());
8730 assert!(pool.to_api_tools().is_empty());
8731 assert!(pool.model_tool_names(&pool.to_api_tools()).is_empty());
8732 assert!(pool.enabled_server_names().is_empty());
8733 assert!(pool.server_names().is_empty());
8734 assert!(pool.connected_servers().is_empty());
8735 assert!(!pool.server_needs_auth("private_a"));
8736 assert!(
8737 pool.authenticate_tool_target("mcp_private_a_authenticate")
8738 .is_none()
8739 );
8740 let (pending, errors) = pool.collect_pending_connects(None);
8741 assert!(pending.is_empty() && errors.is_empty());
8742 assert!(
8743 pool.connect_all().await.is_empty(),
8744 "a denied required server is absent"
8745 );
8746 assert!(
8747 pool.manager_snapshot(Path::new("/unused"), false, &HashMap::new())
8748 .servers
8749 .is_empty()
8750 );
8751 for method in [
8752 "list_mcp_resources",
8753 "list_mcp_resource_templates",
8754 "mcp_read_resource",
8755 "read_mcp_resource",
8756 "mcp_get_prompt",
8757 ] {
8758 let error = pool
8759 .call_tool(
8760 method,
8761 serde_json::json!({"server": "private_a", "uri": "memory://one", "name": "review"}),
8762 )
8763 .await
8764 .unwrap_err();
8765 assert_eq!(
8766 error.to_string(),
8767 "Failed to find MCP server: private_a",
8768 "{method}"
8769 );
8770 }
8771 for method in [
8772 "mcp_private_a_read",
8773 "mcp_private_a_delete",
8774 "mcp_private_a_authenticate",
8775 ] {
8776 assert_eq!(
8777 pool.call_tool(method, serde_json::json!({}))
8778 .await
8779 .unwrap_err()
8780 .to_string(),
8781 format!("Unknown MCP tool name: {method}")
8782 );
8783 }
8784 assert!(pool.begin_authenticate_tool("private_a").await.is_err());
8785 assert!(pool.retry_connection("private_a").await.is_err());
8786 assert!(
8787 pool.add_runtime_server_config("private_a".to_string(), config.clone())
8788 .is_err()
8789 );
8790 let denied = pool
8791 .get_or_connect("private_a")
8792 .await
8793 .err()
8794 .unwrap()
8795 .to_string();
8796 let mut absent = McpPool::new(McpConfig::default());
8797 let missing = absent
8798 .get_or_connect("private_a")
8799 .await
8800 .err()
8801 .unwrap()
8802 .to_string();
8803 assert_eq!(denied, missing);
8804 assert!(sent.lock().unwrap().is_empty(), "no MCP request is sent");
8805 let stale = ceiling_test_connection("private_a", Arc::clone(&sent));
8806 assert!(
8807 pool.store_ready_connection("private_a".to_string(), stale)
8808 .is_err()
8809 );
8810 }
8811
8812 #[tokio::test]
8813 async fn mcp_ceiling_individual_tool_denial_preserves_server_resources_and_sibling_tool() {
8814 let sent = Arc::new(Mutex::new(Vec::new()));
8815 let connection = ceiling_test_connection("private_a", Arc::clone(&sent));
8816 let mut pool = McpPool::new(McpConfig::default())
8817 .with_disallowed_tools(vec!["MCP_PRIVATE_A_DELETE".to_string()]);
8818 pool.connections.insert("private_a".to_string(), connection);
8819 assert_eq!(
8820 pool.all_tools()
8821 .iter()
8822 .map(|(name, _)| name.as_str())
8823 .collect::<Vec<_>>(),
8824 vec!["mcp_private_a_read"]
8825 );
8826 assert_eq!(pool.all_resources().len(), 1);
8827 assert_eq!(pool.all_prompts().len(), 1);
8828 assert!(
8829 pool.call_tool("mcp_private_a_delete", serde_json::json!({}))
8830 .await
8831 .is_err()
8832 );
8833 assert!(sent.lock().unwrap().is_empty());
8834 let resources = pool
8835 .call_tool(
8836 "list_mcp_resources",
8837 serde_json::json!({"server": "private_a"}),
8838 )
8839 .await
8840 .unwrap();
8841 assert_eq!(resources["resources"].as_array().unwrap().len(), 1);
8842 assert_eq!(
8843 pool.call_tool("mcp_private_a_read", serde_json::json!({}))
8844 .await
8845 .unwrap(),
8846 serde_json::json!({"ok": true})
8847 );
8848 assert_eq!(sent.lock().unwrap().len(), 1);
8849 assert_eq!(sent.lock().unwrap()[0]["params"]["name"], "read");
8850 }
8851
8852 #[tokio::test]
8853 async fn mcp_ceiling_child_scoped_meta_calls_do_not_widen_or_mutate_sibling_policy() {
8854 let sent = Arc::new(Mutex::new(Vec::new()));
8855 let private = ceiling_test_connection("private_a", Arc::clone(&sent));
8856 let public = ceiling_test_connection("public", Arc::clone(&sent));
8857 let mut pool = McpPool::new(McpConfig {
8858 servers: HashMap::from([
8859 ("private_a".to_string(), private.config.clone()),
8860 ("public".to_string(), public.config.clone()),
8861 ]),
8862 timeouts: McpTimeouts::default(),
8863 });
8864 pool.connections.insert("private_a".to_string(), private);
8865 pool.connections.insert("public".to_string(), public);
8866 let child_rules = vec!["mcp_private_a_*".to_string()];
8867 for (method, field) in [
8868 ("list_mcp_resources", "resources"),
8869 ("list_mcp_resource_templates", "templates"),
8870 ] {
8871 let child = pool
8872 .call_tool_with_disallowed(method, serde_json::json!({}), &child_rules, None)
8873 .await
8874 .unwrap();
8875 assert_eq!(child[field].as_array().unwrap().len(), 1);
8876 assert_eq!(child[field][0]["server"], "public");
8877 let sibling = pool
8878 .call_tool_with_disallowed(method, serde_json::json!({}), &[], None)
8879 .await
8880 .unwrap();
8881 assert_eq!(sibling[field].as_array().unwrap().len(), 2);
8882 }
8883 assert!(
8884 pool.call_tool_with_disallowed(
8885 "read_mcp_resource",
8886 serde_json::json!({"server":"private_a","uri":"memory://one"}),
8887 &child_rules,
8888 None
8889 )
8890 .await
8891 .is_err()
8892 );
8893 assert!(
8894 pool.call_tool_with_disallowed(
8895 "mcp_private_a_read",
8896 serde_json::json!({}),
8897 &child_rules,
8898 None
8899 )
8900 .await
8901 .is_err()
8902 );
8903 assert!(sent.lock().unwrap().is_empty());
8904 assert_eq!(
8905 pool.call_tool_with_disallowed("mcp_private_a_read", serde_json::json!({}), &[], None)
8906 .await
8907 .unwrap(),
8908 serde_json::json!({"ok":true})
8909 );
8910 }
8911
8912 #[tokio::test]
8913 async fn mcp_ceiling_survives_source_reload_and_blocks_new_runtime_names() {
8914 let dir = tempfile::tempdir().unwrap();
8915 let source = dir.path().join("mcp.json");
8916 fs::write(&source, r#"{"mcpServers": {}}"#).unwrap();
8917 let mut pool = McpPool::from_config_path(&source)
8918 .unwrap()
8919 .with_disallowed_tools(vec!["mcp_private*".to_string()]);
8920 fs::write(&source, r#"{"mcpServers":{"private":{"command":"must-not-execute-private","required":true},"private_a":{"command":"must-not-execute-private"}}}"#).unwrap();
8921 pool.force_reload_config_sources().unwrap();
8922 assert!(pool.connect_all().await.is_empty());
8923 assert!(pool.enabled_server_names().is_empty());
8924 assert!(pool.get_or_connect("private_a").await.is_err());
8925 assert!(
8926 pool.add_runtime_server_config("private_new".to_string(), test_server_config())
8927 .is_err()
8928 );
8929 assert!(!pool.dynamic_servers.read().contains_key("private_new"));
8930 pool.add_runtime_server_config("public".to_string(), test_server_config())
8931 .unwrap();
8932 assert_eq!(pool.enabled_server_names(), vec!["public"]);
8933 }
8934
8935 #[test]
8936 fn mcp_ceiling_namespace_rules_keep_individual_denials_distinct_and_aliases_consistent() {
8937 assert!(McpPool::server_denied_by(&["MCP_A_B_*".to_string()], "a_b"));
8938 assert!(!McpPool::server_denied_by(&["MCP_A_B_*".to_string()], "a"));
8939 assert!(!McpPool::server_denied_by(
8940 &["mcp_a_delete".to_string()],
8941 "a"
8942 ));
8943 assert!(!McpPool::server_denied_by(
8944 &["mcp_a_delete*".to_string()],
8945 "a"
8946 ));
8947 assert!(McpPool::server_denied_by(
8948 &["mcp*".to_string()],
8949 "any_server"
8950 ));
8951 for name in ["mcp_read_resource", "read_mcp_resource"] {
8952 assert!(
8953 McpPool::authorize_call(
8954 &["mcp_a_*".to_string()],
8955 name,
8956 &serde_json::json!({"server":"a"})
8957 )
8958 .is_err()
8959 );
8960 }
8961 }
8962
8963 #[tokio::test]
8964 async fn mcp_ceiling_preserves_ordinary_tool_result_tools_field() {
8965 let sent = Arc::new(Mutex::new(Vec::new()));
8966 let expected = serde_json::json!({"tools":[{"name":"server-owned-data"}], "ok":true});
8967 let mut connection = ceiling_test_connection("public", Arc::clone(&sent));
8968 connection.transport = Box::new(ScriptedValueTransport {
8969 sent,
8970 responses: VecDeque::from([json_frame(
8971 serde_json::json!({"jsonrpc":"2.0","id":1,"result":expected}),
8972 )]),
8973 });
8974 let mut pool = McpPool::new(McpConfig::default());
8975 pool.connections.insert("public".to_string(), connection);
8976 assert_eq!(
8977 pool.call_tool_with_disallowed(
8978 "mcp_public_read",
8979 serde_json::json!({}),
8980 &["mcp_private_*".to_string()],
8981 None
8982 )
8983 .await
8984 .unwrap(),
8985 expected
8986 );
8987 for alias in ["read_mcp_resource", "mcp_read_resource"] {
8988 for denied in ["read_mcp_resource", "mcp_read_resource"] {
8989 assert!(
8990 McpPool::authorize_call(
8991 &[denied.to_string()],
8992 alias,
8993 &serde_json::json!({"server":"public"})
8994 )
8995 .is_err()
8996 );
8997 }
8998 }
8999 }
9000
9001 /// #6213 T7: the resource-URI template check is an authorization decision that
9002 /// runs per URI per advertised template. Pin what it accepts, what it refuses,
9003 /// and that the anchored pattern is compiled once rather than per call.
9004 #[test]
9005 fn resource_uri_template_matching_is_anchored_and_fail_closed() {
9006 // Literal templates are anchored: no suffix may sneak past.
9007 assert!(resource_uri_matches_template(
9008 "file:///readme",
9009 "file:///readme"
9010 ));
9011 assert!(!resource_uri_matches_template(
9012 "file:///readme/extra",
9013 "file:///readme"
9014 ));
9015
9016 // `{id}` is a simple expansion, so it must not cross a path separator.
9017 assert!(resource_uri_matches_template("file:///a", "file:///{id}"));
9018 assert!(!resource_uri_matches_template(
9019 "file:///a/b",
9020 "file:///{id}"
9021 ));
9022
9023 // `{+path}` is a reserved expansion, so it may.
9024 assert!(resource_uri_matches_template(
9025 "file:///a/b/c",
9026 "file:///{+path}"
9027 ));
9028
9029 // An operator this subset does not implement, and a template that never
9030 // closes its expression, both stay uncallable rather than over-matching.
9031 assert!(!resource_uri_matches_template("x", "x{?query}"));
9032 assert!(!resource_uri_matches_template("x", "x{id"));
9033
9034 // The compile happens once per template and is reused.
9035 let first = compiled_resource_template("file:///{path}").expect("template compiles");
9036 let second = compiled_resource_template("file:///{path}").expect("template compiles");
9037 assert!(Arc::ptr_eq(&first, &second));
9038 assert!(compiled_resource_template("x{?query}").is_none());
9039 }
9040
9041 struct DeadTransport;
9042
9043 #[async_trait::async_trait]
9044 impl McpTransport for DeadTransport {
9045 async fn send(&mut self, _msg: Vec<u8>) -> Result<()> {
9046 Ok(())
9047 }
9048
9049 async fn recv(&mut self) -> Result<Vec<u8>> {
9050 Ok(Vec::new())
9051 }
9052
9053 fn probe_dead(&self) -> bool {
9054 true
9055 }
9056 }
9057
9058 fn supervised_pool(name: &str) -> McpPool {
9059 let mut servers = HashMap::new();
9060 servers.insert(name.to_string(), test_server_config());
9061 McpPool::new(McpConfig {
9062 timeouts: McpTimeouts::default(),
9063 servers,
9064 })
9065 }
9066
9067 /// #6187: a dead connection is planned for reconnect, and a failed attempt
9068 /// reports the death once with the diagnosis.
9069 #[test]
9070 fn supervisor_plans_dead_connection_and_reports_failed_reconnect() {
9071 let mut pool = supervised_pool("alpha");
9072 let mut connection = test_connection(Box::new(DeadTransport));
9073 connection.name = "alpha".to_string();
9074 pool.connections.insert("alpha".to_string(), connection);
9075
9076 let plan = pool.plan_supervision();
9077 assert_eq!(plan.due.len(), 1);
9078 assert_eq!(plan.due[0].name, "alpha");
9079 assert!(plan.due[0].fresh_death);
9080 assert!(plan.recovered.is_empty());
9081
9082 let update = pool.resolve_supervision_attempt(
9083 "alpha",
9084 true,
9085 Err(anyhow::anyhow!("connection reset by peer")),
9086 );
9087 assert_eq!(update.died.len(), 1);
9088 assert!(update.died[0].1.contains("connection reset"));
9089 assert!(update.failed.is_empty() && update.recovered.is_empty());
9090
9091 // The failure bought a cooldown: the next sweep attempts nothing and
9092 // reports nothing new.
9093 let plan = pool.plan_supervision();
9094 assert!(plan.due.is_empty());
9095 assert!(plan.recovered.is_empty() && plan.parked.is_empty());
9096 }
9097
9098 /// #6187: recovery is reported on the transition back to alive.
9099 #[test]
9100 fn supervisor_reports_recovery_on_transition() {
9101 let mut pool = supervised_pool("alpha");
9102 let mut dead = test_connection(Box::new(DeadTransport));
9103 dead.name = "alpha".to_string();
9104 pool.connections.insert("alpha".to_string(), dead);
9105
9106 let plan = pool.plan_supervision();
9107 assert_eq!(plan.due.len(), 1);
9108 let update = pool.resolve_supervision_attempt("alpha", true, Err(anyhow::anyhow!("boom")));
9109 assert_eq!(update.died.len(), 1);
9110
9111 // The transport reads alive again (a flapping probe, or a connection
9112 // restored outside the store path): the next sweep reports recovery.
9113 let mut live = test_connection(Box::new(DropCountingTransportForSupervision));
9114 live.name = "alpha".to_string();
9115 pool.connections.insert("alpha".to_string(), live);
9116 let plan = pool.plan_supervision();
9117 assert_eq!(plan.recovered, vec!["alpha".to_string()]);
9118 assert!(plan.due.is_empty());
9119
9120 // Reported once: the sweep after is silent.
9121 let plan = pool.plan_supervision();
9122 assert!(plan.recovered.is_empty() && plan.due.is_empty());
9123 }
9124
9125 struct DropCountingTransportForSupervision;
9126
9127 #[async_trait::async_trait]
9128 impl McpTransport for DropCountingTransportForSupervision {
9129 async fn send(&mut self, _msg: Vec<u8>) -> Result<()> {
9130 Ok(())
9131 }
9132
9133 async fn recv(&mut self) -> Result<Vec<u8>> {
9134 Ok(Vec::new())
9135 }
9136 }
9137
9138 /// #6187: five consecutive failures park the server — no more auto attempts
9139 /// until an explicit retry — and the park is reported once.
9140 #[test]
9141 fn supervisor_parks_after_repeated_failures() {
9142 let mut pool = supervised_pool("alpha");
9143 let mut dead = test_connection(Box::new(DeadTransport));
9144 dead.name = "alpha".to_string();
9145 pool.connections.insert("alpha".to_string(), dead);
9146
9147 for attempt in 0..5 {
9148 let update = pool.resolve_supervision_attempt(
9149 "alpha",
9150 attempt == 0,
9151 Err(anyhow::anyhow!("refused")),
9152 );
9153 if attempt < 4 {
9154 assert!(update.parked.is_empty(), "parks on the fifth failure");
9155 } else {
9156 assert_eq!(update.parked, vec!["alpha".to_string()]);
9157 }
9158 }
9159 // Parked: the plan attempts nothing further.
9160 let plan = pool.plan_supervision();
9161 assert!(plan.due.is_empty());
9162
9163 // A stored-ready connection clears the park.
9164 let mut live = test_connection(Box::new(DropCountingTransportForSupervision));
9165 live.name = "alpha".to_string();
9166 live.catalog_generation = pool.current_catalog_generation();
9167 pool.store_ready_connection("alpha".to_string(), live)
9168 .expect("stores");
9169 assert!(!pool.supervised_parked.contains("alpha"));
9170 assert!(!pool.supervised_dead.contains("alpha"));
9171 }
9172
9173 /// #6187: an explicit retry restarts supervision even when the retry itself
9174 /// fails — the user asked, so the park and dead mark clear.
9175 #[tokio::test]
9176 async fn manual_retry_clears_supervision_marks() {
9177 let mut pool = supervised_pool("alpha");
9178 pool.supervised_dead.insert("alpha".to_string());
9179 pool.supervised_parked.insert("alpha".to_string());
9180 // `mock` is not a real binary, so the retry fails; the marks still clear.
9181 let _ = pool.retry_connection("alpha").await;
9182 assert!(!pool.supervised_dead.contains("alpha"));
9183 assert!(!pool.supervised_parked.contains("alpha"));
9184 }
9185
9186 /// #6187: a dead pipe/socket rebuilds the connection, not just a stale
9187 /// session. Only a refused session id proves the server never ran the
9188 /// request, so only that class may replay a tool call.
9189 #[test]
9190 fn connection_lost_covers_closed_transports_but_only_rejected_sessions_replay() {
9191 use super::wire::{
9192 McpSessionRejected, is_mcp_connection_lost_error, is_mcp_session_rejected_error,
9193 };
9194 for rejected in [
9195 "MCP session expired (transport=sse endpoint=x status=400 Bad Request): session invalid",
9196 "MCP Streamable HTTP session expired; retry with a new session required (404)",
9197 ] {
9198 let err = anyhow::Error::from(McpSessionRejected(rejected.to_string()))
9199 .context("MCP method 'tools/call' failed");
9200 assert!(is_mcp_connection_lost_error(&err), "{rejected}");
9201 assert!(is_mcp_session_rejected_error(&err), "{rejected}");
9202 // The same words without the transport's type are not a refusal:
9203 // a JSON-RPC error answering the request can carry them.
9204 let text_only = anyhow::anyhow!("{rejected}");
9205 assert!(is_mcp_connection_lost_error(&text_only), "{rejected}");
9206 assert!(!is_mcp_session_rejected_error(&text_only), "{rejected}");
9207 }
9208 for ambiguous in [
9209 "MCP session expired: {\"code\":-32000,\"message\":\"session invalid\"}",
9210 "connection reset by peer",
9211 "Stdio transport closed",
9212 "Stdio transport closed (exit status: 1)\nsession invalid",
9213 "SSE transport closed",
9214 "MCP SSE POST send failed (transport=sse endpoint=x): connection closed",
9215 ] {
9216 let err = anyhow::anyhow!("{ambiguous}");
9217 assert!(is_mcp_connection_lost_error(&err), "{ambiguous}");
9218 assert!(!is_mcp_session_rejected_error(&err), "{ambiguous}");
9219 }
9220 let app = anyhow::anyhow!("tool returned an application error");
9221 assert!(!is_mcp_connection_lost_error(&app));
9222 assert!(!is_mcp_session_rejected_error(&app));
9223 }
9224
9225 /// A stdio server that logs every `tools/call` it receives to `$CALL_LOG`,
9226 /// so a test can count how many times the tool really ran. `$FIRST_CALL`
9227 /// picks what happens to the first call after it has run: `exit` kills the
9228 /// child before replying, `session-error` answers it with a JSON-RPC error
9229 /// whose text mentions an invalid session.
9230 #[cfg(unix)]
9231 const COUNTING_STDIO_SERVER: &str = r#"#!/bin/sh
9232 while IFS= read -r line; do
9233 id=$(printf '%s\n' "$line" | sed -n 's/.*"id":"\([^"]*\)".*/\1/p')
9234 case "$line" in
9235 *'"method":"notifications/'*)
9236 ;;
9237 *'"method":"initialize"'*)
9238 printf '{"jsonrpc":"2.0","id":"%s","result":{"protocolVersion":"2024-11-05","serverInfo":{"name":"counter","version":"1.0.0"},"capabilities":{"tools":{}}}}\n' "$id"
9239 ;;
9240 *'"method":"tools/list"'*)
9241 printf '{"jsonrpc":"2.0","id":"%s","result":{"tools":[{"name":"act","inputSchema":{"type":"object"}}]}}\n' "$id"
9242 ;;
9243 *'"method":"tools/call"'*)
9244 echo call >> "$CALL_LOG"
9245 if [ "$(wc -l < "$CALL_LOG")" -eq 1 ]; then
9246 case "$FIRST_CALL" in
9247 exit) exit 0 ;;
9248 session-error)
9249 printf '{"jsonrpc":"2.0","id":"%s","error":{"code":-32000,"message":"session invalid"}}\n' "$id"
9250 continue
9251 ;;
9252 esac
9253 fi
9254 printf '{"jsonrpc":"2.0","id":"%s","result":{"content":[{"type":"text","text":"ok"}]}}\n' "$id"
9255 ;;
9256 *)
9257 [ -n "$id" ] && printf '{"jsonrpc":"2.0","id":"%s","result":{}}\n' "$id"
9258 ;;
9259 esac
9260 done
9261 "#;
9262
9263 #[cfg(unix)]
9264 fn counting_stdio_pool(dir: &Path, first_call: &str) -> (McpPool, PathBuf) {
9265 let script = dir.join("server.sh");
9266 fs::write(&script, COUNTING_STDIO_SERVER).unwrap();
9267 let call_log = dir.join("calls.log");
9268 let mut server = test_server_config();
9269 server.command = Some("sh".to_string());
9270 server.args = vec![script.to_string_lossy().into_owned()];
9271 server.env.insert(
9272 "CALL_LOG".to_string(),
9273 call_log.to_string_lossy().into_owned(),
9274 );
9275 server
9276 .env
9277 .insert("FIRST_CALL".to_string(), first_call.to_string());
9278 server.connect_timeout = Some(10);
9279 server.execute_timeout = Some(10);
9280 let mut cfg = McpConfig::default();
9281 cfg.servers.insert("counter".to_string(), server);
9282 (McpPool::new(cfg), call_log)
9283 }
9284
9285 #[cfg(unix)]
9286 fn server_call_count(call_log: &Path) -> usize {
9287 fs::read_to_string(call_log)
9288 .map(|log| log.lines().count())
9289 .unwrap_or(0)
9290 }
9291
9292 /// A stdio child that dies after reading `tools/call` may have run it: the
9293 /// call fails as unknown, runs once, and the next call gets a fresh child.
9294 #[cfg(unix)]
9295 #[tokio::test]
9296 async fn stdio_child_exit_during_tool_call_is_not_replayed() {
9297 let dir = tempfile::tempdir().unwrap();
9298 let (mut pool, call_log) = counting_stdio_pool(dir.path(), "exit");
9299
9300 let err = pool
9301 .call_tool("mcp_counter_act", serde_json::json!({}))
9302 .await
9303 .expect_err("a call whose child exited mid-flight must not be replayed");
9304 assert!(
9305 format!("{err:#}").contains("outcome unknown, not retried"),
9306 "unexpected error: {err:#}"
9307 );
9308 assert_eq!(server_call_count(&call_log), 1);
9309
9310 let result = pool
9311 .call_tool("mcp_counter_act", serde_json::json!({}))
9312 .await
9313 .expect("the next call reconnects");
9314 assert_eq!(
9315 result,
9316 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
9317 );
9318 assert_eq!(server_call_count(&call_log), 2);
9319 }
9320
9321 /// A JSON-RPC error answering the call id means the server processed the
9322 /// request, even when its text mentions an invalid session: the tool may
9323 /// have acted before failing, so the call is not sent again.
9324 #[cfg(unix)]
9325 #[tokio::test]
9326 async fn json_rpc_session_error_on_tool_call_is_not_replayed() {
9327 let dir = tempfile::tempdir().unwrap();
9328 let (mut pool, call_log) = counting_stdio_pool(dir.path(), "session-error");
9329
9330 let err = pool
9331 .call_tool("mcp_counter_act", serde_json::json!({}))
9332 .await
9333 .expect_err("a session error answering the call must not be replayed");
9334 assert!(
9335 format!("{err:#}").contains("outcome unknown, not retried"),
9336 "unexpected error: {err:#}"
9337 );
9338 assert_eq!(server_call_count(&call_log), 1);
9339
9340 let result = pool
9341 .call_tool("mcp_counter_act", serde_json::json!({}))
9342 .await
9343 .expect("the next call reconnects");
9344 assert_eq!(
9345 result,
9346 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
9347 );
9348 assert_eq!(server_call_count(&call_log), 2);
9349 }
9350
9351 /// A sequential stdio server for the `tools/call` deadline fixtures. Every
9352 /// `initialize` appends `init <pid>` to `$INIT_LOG`, so a test can tell a
9353 /// reused connection from a rebuilt one; every `tools/call` appends
9354 /// `<tool> <id>` to `$CALL_LOG` before it runs, and its reply text is that
9355 /// same line, so a test can tell which request a reply answers. `slow` sleeps
9356 /// `$SLOW_SECS` before replying, `hang` never replies, `fast` replies at once.
9357 #[cfg(unix)]
9358 const DEADLINE_STDIO_SERVER: &str = r#"#!/bin/sh
9359 while IFS= read -r line; do
9360 id=$(printf '%s\n' "$line" | sed -n 's/.*"id":"\([^"]*\)".*/\1/p')
9361 case "$line" in
9362 *'"method":"notifications/'*)
9363 ;;
9364 *'"method":"initialize"'*)
9365 echo "init $$" >> "$INIT_LOG"
9366 printf '{"jsonrpc":"2.0","id":"%s","result":{"protocolVersion":"2024-11-05","serverInfo":{"name":"deadline","version":"1.0.0"},"capabilities":{"tools":{}}}}\n' "$id"
9367 ;;
9368 *'"method":"tools/list"'*)
9369 printf '{"jsonrpc":"2.0","id":"%s","result":{"tools":[{"name":"slow","inputSchema":{"type":"object"}},{"name":"fast","inputSchema":{"type":"object"}},{"name":"hang","inputSchema":{"type":"object"}}]}}\n' "$id"
9370 ;;
9371 *'"method":"tools/call"'*)
9372 tool=$(printf '%s\n' "$line" | sed -n 's/.*"name":"\([^"]*\)".*/\1/p')
9373 echo "$tool $id" >> "$CALL_LOG"
9374 case "$tool" in
9375 slow) sleep "$SLOW_SECS" ;;
9376 hang) while :; do sleep 0.05; done ;;
9377 esac
9378 printf '{"jsonrpc":"2.0","id":"%s","result":{"content":[{"type":"text","text":"%s %s"}]}}\n' "$id" "$tool" "$id"
9379 ;;
9380 *)
9381 [ -n "$id" ] && printf '{"jsonrpc":"2.0","id":"%s","result":{}}\n' "$id"
9382 ;;
9383 esac
9384 done
9385 "#;
9386
9387 /// A pool with one `deadline` stdio server; returns it with the call and
9388 /// init log paths.
9389 #[cfg(unix)]
9390 fn deadline_stdio_pool(
9391 dir: &Path,
9392 slow_secs: &str,
9393 configure: impl FnOnce(&mut McpServerConfig),
9394 ) -> (McpPool, PathBuf, PathBuf) {
9395 let script = dir.join("server.sh");
9396 fs::write(&script, DEADLINE_STDIO_SERVER).unwrap();
9397 let call_log = dir.join("calls.log");
9398 let init_log = dir.join("inits.log");
9399 let mut server = test_server_config();
9400 server.command = Some("sh".to_string());
9401 server.args = vec![script.to_string_lossy().into_owned()];
9402 for (key, value) in [
9403 ("CALL_LOG", call_log.to_string_lossy().into_owned()),
9404 ("INIT_LOG", init_log.to_string_lossy().into_owned()),
9405 ("SLOW_SECS", slow_secs.to_string()),
9406 ] {
9407 server.env.insert(key.to_string(), value);
9408 }
9409 configure(&mut server);
9410 let mut cfg = McpConfig::default();
9411 cfg.servers.insert("deadline".to_string(), server);
9412 (McpPool::new(cfg), call_log, init_log)
9413 }
9414
9415 #[cfg(unix)]
9416 fn log_lines(path: &Path) -> Vec<String> {
9417 fs::read_to_string(path)
9418 .map(|log| log.lines().map(str::to_string).collect())
9419 .unwrap_or_default()
9420 }
9421
9422 /// The result `DEADLINE_STDIO_SERVER` returns for the call it logged as `call`.
9423 #[cfg(unix)]
9424 fn deadline_reply(call: &str) -> serde_json::Value {
9425 serde_json::json!({ "content": [{ "type": "text", "text": call }] })
9426 }
9427
9428 /// #6741 over a real stdio child: a `tools/call` whose reply arrives after
9429 /// the read knob — the receive budget every request used to get — but inside
9430 /// the execute budget completes, and the next call on the same connection gets
9431 /// its own reply. The late frame is consumed once, by the request it answers;
9432 /// nothing is replayed and the child is not restarted.
9433 #[cfg(unix)]
9434 #[tokio::test]
9435 async fn stdio_tool_reply_after_the_read_knob_completes_within_the_execute_budget() {
9436 let dir = tempfile::tempdir().unwrap();
9437 let (mut pool, call_log, init_log) = deadline_stdio_pool(dir.path(), "3", |server| {
9438 server.read_timeout = Some(1);
9439 server.execute_timeout = Some(30);
9440 });
9441
9442 let started = std::time::Instant::now();
9443 let slow = pool
9444 .call_tool("mcp_deadline_slow", serde_json::json!({}))
9445 .await
9446 .expect("a reply inside the execute budget must not be cut off at the read knob");
9447 assert!(
9448 started.elapsed() >= Duration::from_secs(3),
9449 "the reply must really have outlived the 1s read knob"
9450 );
9451 let fast = pool
9452 .call_tool("mcp_deadline_fast", serde_json::json!({}))
9453 .await
9454 .expect("the next call on the same connection must get its own reply");
9455
9456 let calls = log_lines(&call_log);
9457 assert_eq!(calls.len(), 2, "neither call may be replayed: {calls:?}");
9458 assert!(calls[0].starts_with("slow ") && calls[1].starts_with("fast "));
9459 assert_eq!(slow, deadline_reply(&calls[0]));
9460 assert_eq!(fast, deadline_reply(&calls[1]));
9461 assert_eq!(
9462 log_lines(&init_log).len(),
9463 1,
9464 "both calls must share one connection and child"
9465 );
9466 }
9467
9468 /// An explicit `execute_timeout` shorter than the read knob still governs a
9469 /// real stdio `tools/call`: the call fails at its own 2s budget, not at the
9470 /// 120s read knob or the 1800s default. The request is abandoned, not the
9471 /// connection: the child's late reply to it reaches the pipe ahead of the
9472 /// next call's reply and is skipped, so the next call gets its own reply on
9473 /// the same child.
9474 #[cfg(unix)]
9475 #[tokio::test]
9476 async fn explicit_shorter_execute_budget_ends_a_stdio_tool_call_and_skips_its_late_reply() {
9477 let dir = tempfile::tempdir().unwrap();
9478 let (mut pool, call_log, init_log) = deadline_stdio_pool(dir.path(), "3", |server| {
9479 server.execute_timeout = Some(2);
9480 });
9481 pool.get_or_connect("deadline").await.unwrap();
9482
9483 let started = std::time::Instant::now();
9484 let error = pool
9485 .call_tool("mcp_deadline_slow", serde_json::json!({}))
9486 .await
9487 .expect_err("an explicit shorter budget must end the call");
9488 let elapsed = started.elapsed();
9489 assert!(
9490 format!("{error:#}")
9491 .contains("MCP method 'tools/call' on server 'deadline' timed out after 2s"),
9492 "unexpected error: {error:#}"
9493 );
9494 assert!(
9495 elapsed < Duration::from_secs(3),
9496 "the explicit budget, not the 3s reply, must end the call: {elapsed:?}"
9497 );
9498
9499 let fast = pool
9500 .call_tool("mcp_deadline_fast", serde_json::json!({}))
9501 .await
9502 .expect("the kept connection must answer the next call");
9503 let calls = log_lines(&call_log);
9504 assert_eq!(calls.len(), 2, "neither call may be replayed: {calls:?}");
9505 assert!(calls[0].starts_with("slow ") && calls[1].starts_with("fast "));
9506 assert_eq!(
9507 fast,
9508 deadline_reply(&calls[1]),
9509 "the abandoned call's late reply must not answer the next call"
9510 );
9511 assert_eq!(
9512 log_lines(&init_log).len(),
9513 1,
9514 "an expired request must not rebuild the connection"
9515 );
9516 }
9517
9518 /// Owned cancellation of an in-flight stdio `tools/call` under the default
9519 /// 1800s execute budget: cancelling the connection's own token ends the call
9520 /// promptly instead of waiting out the budget, and marks the connection dead.
9521 /// The next call rebuilds it on a fresh child, the cancelled call's child is
9522 /// terminated rather than left running, and the cancelled call is not replayed.
9523 #[cfg(unix)]
9524 #[tokio::test]
9525 async fn cancelling_an_inflight_stdio_tool_call_does_not_wait_for_the_execute_budget() {
9526 let dir = tempfile::tempdir().unwrap();
9527 let (mut pool, call_log, init_log) = deadline_stdio_pool(dir.path(), "0", |_| {});
9528 assert_eq!(
9529 pool.config.servers["deadline"].effective_execute_timeout(&pool.config.timeouts),
9530 1800,
9531 "the fixture must run under the default execute budget"
9532 );
9533 pool.get_or_connect("deadline").await.unwrap();
9534 let cancel = pool.connections["deadline"].cancel_token.clone();
9535
9536 let call = tokio::spawn(async move {
9537 let result = pool
9538 .call_tool("mcp_deadline_hang", serde_json::json!({}))
9539 .await;
9540 (pool, result)
9541 });
9542 tokio::time::timeout(Duration::from_secs(10), async {
9543 while log_lines(&call_log).is_empty() {
9544 tokio::time::sleep(Duration::from_millis(20)).await;
9545 }
9546 })
9547 .await
9548 .expect("the child never received the call");
9549
9550 cancel.cancel();
9551 let (mut pool, result) = tokio::time::timeout(Duration::from_secs(5), call)
9552 .await
9553 .expect("cancellation must not wait for the execute budget")
9554 .expect("the call task panicked");
9555 let error = result.expect_err("a cancelled call must not report success");
9556 assert!(
9557 format!("{error:#}").contains("MCP connection 'deadline' was cancelled"),
9558 "unexpected error: {error:#}"
9559 );
9560 assert!(
9561 pool.connected_servers().is_empty(),
9562 "a cancelled connection must not be reused"
9563 );
9564
9565 let fast = pool
9566 .call_tool("mcp_deadline_fast", serde_json::json!({}))
9567 .await
9568 .expect("the next call must rebuild the connection");
9569 let calls = log_lines(&call_log);
9570 assert_eq!(
9571 calls.len(),
9572 2,
9573 "the cancelled call must not be replayed: {calls:?}"
9574 );
9575 assert!(calls[0].starts_with("hang ") && calls[1].starts_with("fast "));
9576 assert_eq!(fast, deadline_reply(&calls[1]));
9577
9578 let inits = log_lines(&init_log);
9579 assert_eq!(
9580 inits.len(),
9581 2,
9582 "the next call must run on a fresh child: {inits:?}"
9583 );
9584 let cancelled_pid: i32 = inits[0]
9585 .strip_prefix("init ")
9586 .and_then(|pid| pid.parse().ok())
9587 .expect("the init log records the child pid");
9588 tokio::time::timeout(STDIO_SHUTDOWN_GRACE + Duration::from_secs(1), async {
9589 // SAFETY: signal zero only checks whether the pid still exists.
9590 while unsafe { libc::kill(cancelled_pid, 0) } == 0 {
9591 tokio::time::sleep(Duration::from_millis(20)).await;
9592 }
9593 })
9594 .await
9595 .expect("the cancelled call's child must be terminated, not left running");
9596 }
9597
9598 /// A Streamable HTTP server that reads the whole `tools/call` POST and then
9599 /// drops the connection may have run it: the call runs once and the next
9600 /// call succeeds on a rebuilt connection.
9601 #[tokio::test]
9602 async fn streamable_http_reset_after_tool_call_post_is_not_replayed() {
9603 use tokio::io::{AsyncReadExt, AsyncWriteExt};
9604 use tokio::net::TcpListener;
9605
9606 let _lock = lock_mcp_loopback_tests().await;
9607 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
9608 let addr = listener.local_addr().unwrap();
9609 let tool_calls = Arc::new(AtomicUsize::new(0));
9610 let server_tool_calls = Arc::clone(&tool_calls);
9611
9612 let server = tokio::spawn(async move {
9613 loop {
9614 let Ok((mut socket, _)) = listener.accept().await else {
9615 break;
9616 };
9617 let tool_calls = Arc::clone(&server_tool_calls);
9618 tokio::spawn(async move {
9619 let mut request = Vec::new();
9620 let mut buf = [0; 4096];
9621 let header_end = loop {
9622 let Ok(n) = socket.read(&mut buf).await else {
9623 return;
9624 };
9625 if n == 0 {
9626 return;
9627 }
9628 request.extend_from_slice(&buf[..n]);
9629 if let Some(pos) = request.windows(4).position(|w| w == b"\r\n\r\n") {
9630 break pos + 4;
9631 }
9632 };
9633 let headers = String::from_utf8_lossy(&request[..header_end]).to_string();
9634 if headers.starts_with("GET ") {
9635 let _ = socket
9636 .write_all(
9637 b"HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
9638 )
9639 .await;
9640 return;
9641 }
9642 let content_length = headers
9643 .lines()
9644 .find_map(|line| {
9645 let (name, value) = line.split_once(':')?;
9646 name.eq_ignore_ascii_case("content-length")
9647 .then(|| value.trim().parse::<usize>().ok())
9648 .flatten()
9649 })
9650 .unwrap_or(0);
9651 while request.len() < header_end + content_length {
9652 let Ok(n) = socket.read(&mut buf).await else {
9653 return;
9654 };
9655 if n == 0 {
9656 return;
9657 }
9658 request.extend_from_slice(&buf[..n]);
9659 }
9660 let request_json: serde_json::Value =
9661 serde_json::from_slice(&request[header_end..header_end + content_length])
9662 .unwrap();
9663 let method = request_json["method"].as_str().unwrap_or("");
9664 let Some(id) = request_json.get("id").cloned() else {
9665 let _ = socket
9666 .write_all(
9667 b"HTTP/1.1 202 Accepted\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
9668 )
9669 .await;
9670 return;
9671 };
9672 let result = match method {
9673 "initialize" => serde_json::json!({
9674 "protocolVersion": "2024-11-05",
9675 "capabilities": {"tools": {}}
9676 }),
9677 "tools/list" => serde_json::json!({
9678 "tools": [{ "name": "act", "inputSchema": {"type": "object"} }]
9679 }),
9680 "tools/call" => {
9681 // The tool runs, then the connection drops before
9682 // any reply is written.
9683 if tool_calls.fetch_add(1, AtomicOrdering::SeqCst) == 0 {
9684 drop(socket);
9685 return;
9686 }
9687 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
9688 }
9689 _ => serde_json::json!({}),
9690 };
9691 let body =
9692 serde_json::json!({ "jsonrpc": "2.0", "id": id, "result": result }).to_string();
9693 let response = format!(
9694 "HTTP/1.1 200 OK\r\nConnection: close\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}",
9695 body.len(),
9696 body
9697 );
9698 let _ = socket.write_all(response.as_bytes()).await;
9699 let _ = socket.shutdown().await;
9700 });
9701 }
9702 });
9703
9704 let mut server_config = test_server_config();
9705 server_config.command = None;
9706 server_config.url = Some(format!("http://{addr}/mcp"));
9707 server_config.connect_timeout = Some(10);
9708 server_config.execute_timeout = Some(10);
9709 let mut cfg = McpConfig::default();
9710 cfg.servers.insert("remote".to_string(), server_config);
9711 let mut pool = McpPool::new(cfg);
9712
9713 let err = pool
9714 .call_tool("mcp_remote_act", serde_json::json!({}))
9715 .await
9716 .expect_err("a call whose connection dropped mid-flight must not be replayed");
9717 assert!(
9718 format!("{err:#}").contains("outcome unknown, not retried"),
9719 "unexpected error: {err:#}"
9720 );
9721 assert_eq!(tool_calls.load(AtomicOrdering::SeqCst), 1);
9722
9723 let result = pool
9724 .call_tool("mcp_remote_act", serde_json::json!({}))
9725 .await
9726 .expect("the next call reconnects");
9727 assert_eq!(
9728 result,
9729 serde_json::json!({ "content": [{ "type": "text", "text": "ok" }] })
9730 );
9731 assert_eq!(tool_calls.load(AtomicOrdering::SeqCst), 2);
9732
9733 server.abort();
9734 }
9735
9736 // Executed both as an ordinary no-op test and as an isolated OS-process worker.
9737 #[test]
9738 fn mcp_transaction_child_worker() {
9739 let Some(path) = std::env::var_os("CW_MCP_TRANSACTION_TEST_PATH") else {
9740 return;
9741 };
9742 let path = PathBuf::from(path);
9743 let mode = std::env::var("CW_MCP_TRANSACTION_TEST_MODE").unwrap();
9744 if mode == "init" {
9745 init_config(&path, false).unwrap();
9746 return;
9747 }
9748 mutate_config(&path, None, |cfg| {
9749 if mode == "hold" {
9750 fs::write(path.with_extension("entered"), b"ready")?;
9751 let deadline = std::time::Instant::now() + Duration::from_secs(10);
9752 while !path.with_extension("release").exists() {
9753 anyhow::ensure!(
9754 std::time::Instant::now() < deadline,
9755 "fixture release timed out"
9756 );
9757 std::thread::sleep(Duration::from_millis(5));
9758 }
9759 }
9760 let server: McpServerConfig =
9761 serde_json::from_value(serde_json::json!({"command":"fixture-command"}))?;
9762 cfg.servers.insert(mode.clone(), server);
9763 Ok(())
9764 })
9765 .unwrap();
9766 }
9767
9768 fn mcp_transaction_spawn_worker(path: &Path, mode: &str) -> std::process::Child {
9769 std::process::Command::new(std::env::current_exe().unwrap())
9770 .args([
9771 "--exact",
9772 "mcp::tests::mcp_transaction_child_worker",
9773 "--nocapture",
9774 ])
9775 .env("CW_MCP_TRANSACTION_TEST_PATH", path)
9776 .env("CW_MCP_TRANSACTION_TEST_MODE", mode)
9777 .stdin(std::process::Stdio::null())
9778 .stdout(std::process::Stdio::piped())
9779 .stderr(std::process::Stdio::piped())
9780 .spawn()
9781 .unwrap()
9782 }
9783
9784 #[test]
9785 fn mcp_transaction_independent_process_writers_and_init_preserve_updates() {
9786 for mode in ["second", "init"] {
9787 let root = tempfile::tempdir().unwrap();
9788 let path = root.path().join("mcp.json");
9789 let first = mcp_transaction_spawn_worker(&path, "hold");
9790 let deadline = std::time::Instant::now() + Duration::from_secs(5);
9791 while !path.with_extension("entered").exists() {
9792 assert!(
9793 std::time::Instant::now() < deadline,
9794 "first worker did not acquire lock"
9795 );
9796 std::thread::sleep(Duration::from_millis(5));
9797 }
9798 // On Unix the competing process uses an alias of the same directory.
9799 #[cfg(unix)]
9800 let other_path = {
9801 let alias = root.path().join("alias");
9802 std::os::unix::fs::symlink(root.path(), &alias).unwrap();
9803 alias.join("mcp.json")
9804 };
9805 #[cfg(not(unix))]
9806 let other_path = path.clone();
9807 let mut second = mcp_transaction_spawn_worker(&other_path, mode);
9808 std::thread::sleep(Duration::from_millis(40));
9809 assert!(
9810 second.try_wait().unwrap().is_none(),
9811 "second writer must wait for shared lock"
9812 );
9813 fs::write(path.with_extension("release"), b"go").unwrap();
9814 for child in [first, second] {
9815 let result = child.wait_with_output().unwrap();
9816 assert!(
9817 result.status.success(),
9818 "{}",
9819 String::from_utf8_lossy(&result.stderr)
9820 );
9821 }
9822 let cfg = load_config(&path).unwrap();
9823 assert!(cfg.servers.contains_key("hold"));
9824 if mode == "second" {
9825 assert!(cfg.servers.contains_key("second"));
9826 }
9827 assert!(
9828 !cfg.servers.contains_key("example"),
9829 "init must not overwrite a concurrent add"
9830 );
9831 }
9832 }
9833
9834 #[test]
9835 fn mcp_transaction_preserves_unknown_fields_alias_and_rejects_stale_revision() {
9836 let root = tempfile::tempdir().unwrap();
9837 let path = root.path().join("mcp.json");
9838 fs::write(&path, r#"{"owner_note":{"keep":true},"timeouts":{"connect_timeout":10,"custom":42},"mcpServers":{"one":{"command":"one","extension":{"keep":1}}}}"#).unwrap();
9839 let before = read_config_revision(&path).unwrap();
9840 let (_, after) = mutate_config(&path, Some(&before), |cfg| {
9841 cfg.servers.get_mut("one").unwrap().enabled = false;
9842 Ok(())
9843 })
9844 .unwrap();
9845 assert_ne!(before, after);
9846 let raw: serde_json::Value = serde_json::from_slice(&fs::read(&path).unwrap()).unwrap();
9847 assert_eq!(raw["owner_note"]["keep"], true);
9848 assert_eq!(raw["timeouts"]["custom"], 42);
9849 assert_eq!(raw["mcpServers"]["one"]["extension"]["keep"], 1);
9850 assert!(raw.get("servers").is_none());
9851 let bytes = fs::read(&path).unwrap();
9852 let err = mutate_config(&path, Some(&before), |cfg| {
9853 cfg.servers.clear();
9854 Ok(())
9855 })
9856 .unwrap_err();
9857 assert!(err.is::<McpRevisionConflict>());
9858 assert_eq!(fs::read(&path).unwrap(), bytes);
9859 assert_eq!(
9860 mutate_config(&path, Some(&after), |_| Ok(())).unwrap().1,
9861 after
9862 );
9863 assert_eq!(
9864 fs::read(&path).unwrap(),
9865 bytes,
9866 "no-op must not rewrite the document"
9867 );
9868 }
9869
9870 #[test]
9871 fn mcp_transaction_fails_closed_for_malformed_document_and_symlink() {
9872 let root = tempfile::tempdir().unwrap();
9873 let path = root.path().join("mcp.json");
9874 fs::write(&path, "{private-malformed-fixture").unwrap();
9875 let err = mutate_config(&path, None, |_| Ok(())).unwrap_err();
9876 assert!(!err.to_string().contains("private-malformed-fixture"));
9877 assert!(init_config(&path, true).is_err());
9878 assert_eq!(
9879 fs::read_to_string(&path).unwrap(),
9880 "{private-malformed-fixture"
9881 );
9882 #[cfg(unix)]
9883 {
9884 let link = root.path().join("linked.json");
9885 std::os::unix::fs::symlink(&path, &link).unwrap();
9886 assert!(mutate_config(&link, None, |_| Ok(())).is_err());
9887 assert!(init_config(&link, true).is_err());
9888 }
9889 }
9890
9891 #[test]
9892 fn only_a_reviewed_plugin_read_only_hint_relaxes_approval() {
9893 let tool: McpTool = serde_json::from_value(serde_json::json!({
9894 "name": "page_snapshot",
9895 "inputSchema": {"type": "object"},
9896 "annotations": {"readOnlyHint": true, "destructiveHint": false}
9897 }))
9898 .expect("annotated tool parses");
9899 assert_eq!(
9900 approval_hint_for(&tool, true),
9901 Some(McpToolApprovalHint::TrustedReadOnly)
9902 );
9903 // The same claim from a server no plugin review covers is not trusted.
9904 assert_eq!(approval_hint_for(&tool, false), None);
9905
9906 let destructive: McpTool = serde_json::from_value(serde_json::json!({
9907 "name": "delete_rows",
9908 "annotations": {"readOnlyHint": true, "destructiveHint": true}
9909 }))
9910 .expect("annotated tool parses");
9911 // A tool that claims both keeps its prompt, from any server.
9912 assert_eq!(
9913 approval_hint_for(&destructive, true),
9914 Some(McpToolApprovalHint::Destructive)
9915 );
9916 assert_eq!(
9917 approval_hint_for(&destructive, false),
9918 Some(McpToolApprovalHint::Destructive)
9919 );
9920
9921 let bare: McpTool = serde_json::from_value(serde_json::json!({"name": "echo"}))
9922 .expect("unannotated tool parses");
9923 assert_eq!(approval_hint_for(&bare, true), None);
9924 }
9925
9926 fn stdio_server(args: Vec<String>) -> McpServerConfig {
9927 serde_json::from_value(serde_json::json!({ "command": "node", "args": args }))
9928 .expect("stdio server config")
9929 }
9930
9931 #[test]
9932 fn user_server_launching_the_computer_use_bundle_is_recognized() {
9933 let dir = tempfile::tempdir().expect("tempdir");
9934 let bundle = dir
9935 .path()
9936 .join("Codewhale Computer Use.app/Contents/Resources/plugin");
9937 std::fs::create_dir_all(bundle.join("mcp")).expect("bundle dirs");
9938 std::fs::write(bundle.join("mcp/server.mjs"), "").expect("server");
9939 std::fs::write(bundle.join("plugin.json"), r#"{"name": "computer-use"}"#).expect("manifest");
9940 let script = bundle.join("mcp/server.mjs").to_string_lossy().to_string();
9941 assert_eq!(
9942 launches_computer_use_plugin(&stdio_server(vec![script.clone()])),
9943 Some(script)
9944 );
9945
9946 // Another plugin's server with the same layout is not a duplicate.
9947 let other = dir.path().join("other-plugin");
9948 std::fs::create_dir_all(other.join("mcp")).expect("other dirs");
9949 std::fs::write(other.join("plugin.json"), r#"{"name": "browser-tools"}"#).expect("manifest");
9950 let other_script = other.join("mcp/server.mjs").to_string_lossy().to_string();
9951 assert_eq!(
9952 launches_computer_use_plugin(&stdio_server(vec![other_script])),
9953 None
9954 );
9955
9956 // Without a readable manifest the bundle's path shape still counts.
9957 let shaped = "/opt/computer-use/mcp/server.mjs".to_string();
9958 assert_eq!(
9959 launches_computer_use_plugin(&stdio_server(vec![shaped.clone()])),
9960 Some(shaped)
9961 );
9962 assert_eq!(
9963 launches_computer_use_plugin(&stdio_server(vec!["/opt/tools/mcp/server.mjs".to_string()])),
9964 None
9965 );
9966 }
9967
9968 #[test]
9969 fn computer_use_duplicate_warning_needs_the_builtin_bundle_enabled() {
9970 // A user entry alone (no enabled built-in bundle) is never flagged.
9971 let mut config = McpConfig::default();
9972 config.servers.insert(
9973 "codewhale-cu".to_string(),
9974 stdio_server(vec!["/opt/computer-use/mcp/server.mjs".to_string()]),
9975 );
9976 assert!(duplicate_computer_use_servers(&config).is_empty());
9977 }
9978
9979 #[test]
9980 fn computer_use_duplicate_warning_names_user_copies_of_the_enabled_bundle() {
9981 let dir = tempfile::tempdir().expect("tempdir");
9982 let plugin_base = dir.path().join("plugins/computer-use");
9983 fs::create_dir_all(plugin_base.join("mcp")).expect("plugin dirs");
9984 fs::write(
9985 plugin_base.join("plugin.toml"),
9986 "schema_version = 1\n[plugin]\nname = \"computer-use\"\nversion = \"1.0.0\"\n",
9987 )
9988 .expect("plugin manifest");
9989 let (_, authority) = active_plugin_fixture(&plugin_base);
9990 let mut builtin = stdio_server(vec![
9991 plugin_base
9992 .join("mcp/server.mjs")
9993 .to_string_lossy()
9994 .to_string(),
9995 ]);
9996 builtin.reviewed_plugin = Some(
9997 ReviewedPluginMcpSource::from_authority(
9998 authority,
9999 None,
10000 Arc::new(crate::plugins::HostEnvironment::default()),
10001 )
10002 .expect("reviewed source"),
10003 );
10004
10005 let mut config = McpConfig::default();
10006 config
10007 .servers
10008 .insert("plugin-computer-use".to_string(), builtin);
10009 config.servers.insert(
10010 "codewhale-cu".to_string(),
10011 stdio_server(vec!["/opt/computer-use/mcp/server.mjs".to_string()]),
10012 );
10013 config.servers.insert(
10014 "browser-tools".to_string(),
10015 stdio_server(vec!["/opt/tools/mcp/server.mjs".to_string()]),
10016 );
10017 let mut disabled_copy = stdio_server(vec!["/srv/computer_use/mcp/server.mjs".to_string()]);
10018 disabled_copy.enabled = false;
10019 config.servers.insert("old-cu".to_string(), disabled_copy);
10020
10021 // Only the enabled user copy is named, with the argument that gave it
10022 // away; the bundle itself, other plugins and disabled entries are not.
10023 assert_eq!(
10024 duplicate_computer_use_servers(&config),
10025 vec![(
10026 "codewhale-cu".to_string(),
10027 "/opt/computer-use/mcp/server.mjs".to_string()
10028 )]
10029 );
10030
10031 // Disabling the built-in bundle removes the warning.
10032 config
10033 .servers
10034 .get_mut("plugin-computer-use")
10035 .expect("bundle entry")
10036 .enabled = false;
10037 assert!(duplicate_computer_use_servers(&config).is_empty());
10038 }
10039
10040 /// Founder run: an `uvx mcp-proxy-for-aws` server answered `initialize` with
10041 /// JSON-RPC -32602 and the only thing surfaced was
10042 /// "MCP error in 'initialize': {...}". The proxy had explained itself on
10043 /// stderr (expired AWS login). The error now says the server rejected the
10044 /// handshake, names what was launched, and carries that last stderr line.
10045 #[cfg(unix)]
10046 #[tokio::test]
10047 async fn initialize_rejection_names_the_server_command_and_its_stderr_reason() {
10048 let mut config = test_server_config();
10049 config.command = Some("sh".to_string());
10050 config.args = vec![
10051 "-c".to_string(),
10052 concat!(
10053 "read line; ",
10054 "echo 'LoginRefreshRequired: Please reauthenticate using aws login' 1>&2; ",
10055 "echo '{\"jsonrpc\":\"2.0\",\"id\":\"1\",\"error\":{\"code\":-32602,\"message\":\"Invalid request parameters\"}}'; ",
10056 "sleep 5"
10057 )
10058 .to_string(),
10059 ];
10060 let error = McpConnection::connect_with_policy(
10061 "aws".to_string(),
10062 config,
10063 &McpTimeouts::default(),
10064 None,
10065 )
10066 .await
10067 .err()
10068 .expect("a JSON-RPC error on initialize ends the handshake");
10069 let text = format_mcp_error_for_display(&error);
10070 assert!(
10071 text.contains("MCP server 'aws' rejected initialize"),
10072 "{text}"
10073 );
10074 assert!(text.contains("command `sh`"), "{text}");
10075 assert!(text.contains("-32602"), "{text}");
10076 assert!(
10077 text.contains("server stderr: LoginRefreshRequired: Please reauthenticate using aws login"),
10078 "{text}"
10079 );
10080 // The founder's exact stderr is an AWS CLI `aws login` session, renewed
10081 // with `aws login` (not `aws sso login`).
10082 assert!(
10083 text.contains(
10084 "AWS credentials expired: run `aws login` in a terminal, then `/mcp retry aws`"
10085 ),
10086 "{text}"
10087 );
10088 }
10089
10090 /// The file the founder had carried both `"disabled": true` and
10091 /// `"enabled": false`; a hand edit can leave them disagreeing. Enabling (and
10092 /// disabling) must always write the pair so the two keys agree.
10093 #[test]
10094 fn set_server_enabled_makes_the_enabled_and_disabled_keys_agree() {
10095 let dir = tempfile::tempdir().unwrap();
10096 let path = dir.path().join("mcp.json");
10097 fs::write(
10098 &path,
10099 r#"{"servers":{"linear":{"url":"https://mcp.linear.app/mcp","enabled":true,"disabled":true}}}"#,
10100 )
10101 .unwrap();
10102 assert!(!load_config(&path).unwrap().servers["linear"].is_enabled());
10103
10104 set_server_enabled(&path, "linear", true).unwrap();
10105 let raw: serde_json::Value = serde_json::from_str(&fs::read_to_string(&path).unwrap()).unwrap();
10106 assert_eq!(raw["servers"]["linear"]["enabled"], serde_json::json!(true));
10107 assert_eq!(
10108 raw["servers"]["linear"]["disabled"],
10109 serde_json::json!(false)
10110 );
10111 assert!(load_config(&path).unwrap().servers["linear"].is_enabled());
10112
10113 set_server_enabled(&path, "linear", false).unwrap();
10114 let raw: serde_json::Value = serde_json::from_str(&fs::read_to_string(&path).unwrap()).unwrap();
10115 assert_eq!(
10116 raw["servers"]["linear"]["enabled"],
10117 serde_json::json!(false)
10118 );
10119 assert_eq!(
10120 raw["servers"]["linear"]["disabled"],
10121 serde_json::json!(true)
10122 );
10123 }
10124
10125 /// Lazy boot (#6033) leaves an unselected server with no connection, no
10126 /// failure and no observed capabilities. That is "never started" — its
10127 /// recovery is `connect`, not `reconnect`, and an OAuth-capable one is not
10128 /// yet known to need a login.
10129 #[test]
10130 fn never_started_server_recovers_with_connect_not_reconnect() {
10131 let snapshot = McpServerSnapshot {
10132 name: "lazy".into(),
10133 enabled: true,
10134 required: false,
10135 transport: "stdio".into(),
10136 command_or_url: "lazy-mcp".into(),
10137 connect_timeout: 5,
10138 execute_timeout: 5,
10139 read_timeout: 5,
10140 connected: false,
10141 error: None,
10142 auth_required: false,
10143 capability_metadata: McpServerCapabilityMetadata::NotObserved,
10144 tools: Vec::new(),
10145 resources: Vec::new(),
10146 prompts: Vec::new(),
10147 };
10148 assert!(!snapshot.started());
10149 assert_eq!(
10150 snapshot.recovery_kind(false),
10151 Some(McpRecoveryKind::Connect)
10152 );
10153 assert_eq!(snapshot.recovery_kind(true), Some(McpRecoveryKind::Connect));
10154
10155 // A server that had a live connection and lost it is a reconnect.
10156 let dropped = McpServerSnapshot {
10157 capability_metadata: McpServerCapabilityMetadata::LegacyFallback,
10158 ..snapshot
10159 };
10160 assert!(dropped.started());
10161 assert_eq!(
10162 dropped.recovery_kind(false),
10163 Some(McpRecoveryKind::Reconnect)
10164 );
10165 }
10166
10167 #[tokio::test]
10168 async fn initialize_captures_sanitized_server_instructions() {
10169 let transport = ScriptedValueTransport {
10170 sent: Arc::new(Mutex::new(Vec::new())),
10171 responses: VecDeque::from([json_frame(serde_json::json!({
10172 "jsonrpc": "2.0",
10173 "id": 1,
10174 "result": {
10175 "protocolVersion": "2024-11-05",
10176 "serverInfo": {"name": "guided", "version": "1.0.0"},
10177 "capabilities": {"tools": {}},
10178 "instructions": " Prefer search\u{7} before fetch.\r\n\tKeep queries short.\u{202E} "
10179 }
10180 }))]),
10181 };
10182 let mut conn = test_connection(Box::new(transport));
10183 conn.initialize().await.expect("initialize");
10184 assert_eq!(
10185 conn.instructions(),
10186 Some("Prefer search before fetch.\n\tKeep queries short.")
10187 );
10188
10189 let transport = ScriptedValueTransport {
10190 sent: Arc::new(Mutex::new(Vec::new())),
10191 responses: VecDeque::from([json_frame(serde_json::json!({
10192 "jsonrpc": "2.0",
10193 "id": 1,
10194 "result": {
10195 "protocolVersion": "2024-11-05",
10196 "serverInfo": {"name": "verbose", "version": "1.0.0"},
10197 "capabilities": {"tools": {}},
10198 "instructions": "x".repeat(10_000)
10199 }
10200 }))]),
10201 };
10202 let mut conn = test_connection(Box::new(transport));
10203 conn.initialize().await.expect("initialize");
10204 let capped = conn.instructions().expect("capped guidance");
10205 assert!(capped.len() <= codewhale_mcp::MAX_SERVER_INSTRUCTIONS_BYTES);
10206 assert!(capped.ends_with("[truncated]"));
10207 }
10208
10209 #[tokio::test]
10210 async fn initialize_ignores_non_string_instructions_without_failing_the_handshake() {
10211 let transport = ScriptedValueTransport {
10212 sent: Arc::new(Mutex::new(Vec::new())),
10213 responses: VecDeque::from([json_frame(serde_json::json!({
10214 "jsonrpc": "2.0",
10215 "id": 1,
10216 "result": {
10217 "protocolVersion": "2024-11-05",
10218 "serverInfo": {"name": "odd", "version": "1.0.0"},
10219 "capabilities": {"tools": {}},
10220 "instructions": {"text": "not a string"}
10221 }
10222 }))]),
10223 };
10224 let mut conn = test_connection(Box::new(transport));
10225 conn.initialize()
10226 .await
10227 .expect("non-string instructions must not fail the handshake");
10228 assert_eq!(conn.instructions(), None);
10229 }
10230
10231 #[tokio::test]
10232 async fn mcp_server_instructions_exclude_servers_whose_tools_are_all_denied() {
10233 let mut pool = McpPool::new(McpConfig::default())
10234 .with_disallowed_tools(vec!["mcp_denied_write".to_string()]);
10235 pool.insert_test_connection("guided", &["search"], Some("Use search."));
10236 pool.insert_test_connection("denied", &["write"], Some("Denied guidance."));
10237 pool.insert_test_connection("hidden", &["read"], Some("Hidden guidance."));
10238 pool.insert_test_connection("silent", &["ping"], None);
10239
10240 // `hidden` is allowed by the pool but absent from the turn's catalog.
10241 let visible = |name: &str| name != "mcp_hidden_read";
10242 assert_eq!(
10243 pool.model_server_instructions(visible),
10244 vec![("guided".to_string(), "Use search.".to_string())]
10245 );
10246 }
10247
10248 #[tokio::test]
10249 async fn computer_use_call_carries_an_attested_decision_only_when_a_person_approved() {
10250 let respond = |id: u64| {
10251 json_frame(serde_json::json!({"jsonrpc": "2.0", "id": id, "result": {"ok": true}}))
10252 };
10253 let sent = Arc::new(Mutex::new(Vec::new()));
10254 let mut connection = test_connection(Box::new(ScriptedValueTransport {
10255 sent: Arc::clone(&sent),
10256 responses: VecDeque::from([respond(1), respond(2), respond(3)]),
10257 }));
10258 let key = [7_u8; 32];
10259 connection.decision_key = Some(key);
10260 let args = serde_json::json!({"action": "allow", "app": "Safari", "remember": 1.0});
10261 let decision = crate::core::engine::HumanDecision::for_test("mcp_fixture_consent", &args);
10262 connection
10263 .call_tool_decided("consent", args.clone(), 5, Some(&decision))
10264 .await
10265 .expect("decided call");
10266 connection
10267 .call_tool_decided("consent", args.clone(), 5, None)
10268 .await
10269 .expect("plain call");
10270 connection.decision_key = None;
10271 connection
10272 .call_tool_decided("consent", args.clone(), 5, Some(&decision))
10273 .await
10274 .expect("call without a shared key");
10275
10276 let sent = sent.lock().unwrap().clone();
10277 let attested = &sent[0]["params"]["_meta"][COMPUTER_USE_DECISION_META];
10278 let nonce = attested["nonce"].as_str().expect("nonce");
10279 let args_json = attested["args_json"].as_str().expect("args_json");
10280 assert_eq!(
10281 serde_json::from_str::<serde_json::Value>(args_json).unwrap(),
10282 sent[0]["params"]["arguments"]
10283 );
10284 let message = [
10285 b"consent".as_slice(),
10286 b"\0",
10287 args_json.as_bytes(),
10288 b"\0",
10289 nonce.as_bytes(),
10290 ]
10291 .concat();
10292 let expected = ring::hmac::sign(
10293 &ring::hmac::Key::new(ring::hmac::HMAC_SHA256, &key),
10294 &message,
10295 );
10296 assert_eq!(attested["mac"], hex_encode(expected.as_ref()));
10297 assert!(sent[1]["params"].get("_meta").is_none(), "{}", sent[1]);
10298 assert!(sent[2]["params"].get("_meta").is_none(), "{}", sent[2]);
10299 }
10300
10301 #[tokio::test]
10302 async fn copied_human_decision_cannot_authorize_another_mcp_call() {
10303 let input = serde_json::json!({"action": "allow", "app": "Safari"});
10304 let decision = crate::core::engine::HumanDecision::for_test("mcp_fixture_consent", &input);
10305 let mut pool = McpPool::new(McpConfig::default());
10306 for (name, arguments) in [
10307 (
10308 "mcp_fixture_consent",
10309 serde_json::json!({"action": "allow", "app": "Terminal"}),
10310 ),
10311 ("mcp_other_consent", input.clone()),
10312 ] {
10313 let error = pool
10314 .call_tool_with_decision(name, arguments, Some(&decision))
10315 .await
10316 .unwrap_err();
10317 assert!(
10318 error
10319 .to_string()
10320 .contains("does not authorize this exact MCP call"),
10321 "{error:#}"
10322 );
10323 }
10324 assert!(decision.authorizes("mcp_fixture_consent", &input));
10325 }
10326
10327 /// Actual shipped plugin, reviewed and staged by the normal plugin authority.
10328 /// Every backend and state path is isolated; this fixture never drives a real app.
10329 pub(crate) fn computer_use_test_fixture() -> (
10330 tempfile::TempDir,
10331 Arc<crate::plugins::PluginRegistry>,
10332 McpPool,
10333 String,
10334 ) {
10335 let root = tempfile::tempdir().expect("private fixture root");
10336 let workspace = root.path().join("workspace");
10337 fs::create_dir(&workspace).unwrap();
10338 let plugin_base = root.path().join("plugins/computer-use");
10339 let source = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("plugins/computer-use");
10340 let mut pending = vec![(source, plugin_base.clone())];
10341 while let Some((from, to)) = pending.pop() {
10342 fs::create_dir_all(&to).unwrap();
10343 for entry in fs::read_dir(from).unwrap() {
10344 let entry = entry.unwrap();
10345 let destination = to.join(entry.file_name());
10346 if entry.file_type().unwrap().is_dir() {
10347 pending.push((entry.path(), destination));
10348 } else {
10349 fs::copy(entry.path(), destination).unwrap();
10350 }
10351 }
10352 }
10353 let state = root.path().join("cu-state");
10354 let recordings = root.path().join("recordings");
10355 let fake = plugin_base.join("tests/fixtures/fake-backend.mjs");
10356 fs::write(
10357 plugin_base.join("plugin.toml"),
10358 format!(
10359 "schema_version = 1\n[plugin]\nname = \"computer-use\"\nversion = \"1.0.0\"\n\
10360 [mcp_servers.local]\ncommand = \"node\"\nargs = [\"mcp/server.mjs\"]\nconnect_timeout = 5\n\
10361 [mcp_servers.local.env]\nCODEWHALE_CU_APP = \"off\"\n\
10362 CODEWHALE_CU_STATE_DIR = {}\nCODEWHALE_CU_RECORDINGS_DIR = {}\nCODEWHALE_CU_TEST_BACKEND = {}\n",
10363 serde_json::to_string(&state.to_string_lossy()).unwrap(),
10364 serde_json::to_string(&recordings.to_string_lossy()).unwrap(),
10365 serde_json::to_string(&fake.to_string_lossy()).unwrap(),
10366 ),
10367 )
10368 .unwrap();
10369 let discovery = crate::plugins::discovery::DiscoveryConfig {
10370 workspace,
10371 user_plugins_dir: root.path().join("plugins"),
10372 workspace_plugins_dir: root.path().join("unused-workspace-plugins"),
10373 builtin_plugin_dirs: Vec::new(),
10374 state_path: root.path().join("plugin-state/state.json"),
10375 };
10376 let mut registry = crate::plugins::discovery::discover_with_config(&discovery);
10377 registry.trust("computer-use").unwrap();
10378 registry.enable("computer-use").unwrap();
10379 let plugin = registry.get("computer-use").unwrap().clone();
10380 let authority = registry.authority_for("computer-use").unwrap();
10381 let config = merge_plugin_mcp_servers_from_plugins(
10382 McpConfig::default(),
10383 vec![("computer-use".to_string(), plugin, authority)],
10384 )
10385 .unwrap();
10386 let server = config
10387 .servers
10388 .keys()
10389 .next()
10390 .expect("plugin MCP server")
10391 .clone();
10392 (root, Arc::new(registry), McpPool::new(config), server)
10393 }
10394
10395 #[tokio::test]
10396 async fn computer_use_real_plugin_host_handshake_rejects_tamper_replay_and_late_keys() {
10397 // The bundled Computer Use plugin applies only to macOS hosts; elsewhere
10398 // there is no live plugin to start.
10399 if !cfg!(target_os = "macos") {
10400 return;
10401 }
10402 let _env = crate::test_support::lock_test_env();
10403 let (root, _registry, mut pool, server) = computer_use_test_fixture();
10404 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", root.path());
10405 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
10406 let connection = pool
10407 .get_or_connect(&server)
10408 .await
10409 .expect("actual reviewed Node plugin");
10410 let key = connection
10411 .decision_key
10412 .expect("first-message key from the host");
10413 let args = serde_json::json!({"scope":"foreground", "remember":false});
10414 let decision = crate::core::engine::HumanDecision::for_test("mcp_fixture_consent_allow", &args);
10415 let decode = |result: serde_json::Value| {
10416 serde_json::from_str::<serde_json::Value>(result["content"][0]["text"].as_str().unwrap())
10417 .unwrap()
10418 };
10419 let unsigned = decode(
10420 connection
10421 .call_tool_decided("consent_allow", args.clone(), 5, None)
10422 .await
10423 .unwrap(),
10424 );
10425 assert_eq!(
10426 unsigned["error"]["code"], "consent_needs_user",
10427 "{unsigned}"
10428 );
10429 let approved = decode(
10430 connection
10431 .call_tool_decided("consent_allow", args.clone(), 5, Some(&decision))
10432 .await
10433 .unwrap(),
10434 );
10435 assert_eq!(approved["ok"], true, "{approved}");
10436
10437 // A later key notification must not replace the first connection key.
10438 let wrong_key = [8_u8; 32];
10439 connection.send(serde_json::json!({"jsonrpc":"2.0","method":COMPUTER_USE_HOST_KEYS_METHOD,"params":{"decision_key":hex_encode(&wrong_key)}})).await.unwrap();
10440 let fresh = attest_decision(&key, "consent_revoke", &args).unwrap();
10441 for (tool, arguments, tag) in [
10442 (
10443 "consent_revoke",
10444 args.clone(),
10445 attest_decision(&wrong_key, "consent_revoke", &args).unwrap(),
10446 ),
10447 (
10448 "consent_revoke",
10449 serde_json::json!({"scope":"foreground","remember":true}),
10450 fresh.clone(),
10451 ),
10452 ("consent_allow", args.clone(), fresh.clone()),
10453 ] {
10454 let rejected = decode(connection.call_method("tools/call", serde_json::json!({"name":tool,"arguments":arguments,"_meta":{COMPUTER_USE_DECISION_META:tag}}), 5).await.unwrap());
10455 assert_eq!(
10456 rejected["error"]["code"], "consent_needs_user",
10457 "{rejected}"
10458 );
10459 }
10460 let exact = serde_json::json!({"name":"consent_revoke","arguments":args,"_meta":{COMPUTER_USE_DECISION_META:fresh}});
10461 let allowed = decode(
10462 connection
10463 .call_method("tools/call", exact.clone(), 5)
10464 .await
10465 .unwrap(),
10466 );
10467 assert_eq!(
10468 allowed["ok"], true,
10469 "the original key remains authoritative: {allowed}"
10470 );
10471 let replay = decode(
10472 connection
10473 .call_method("tools/call", exact, 5)
10474 .await
10475 .unwrap(),
10476 );
10477 assert_eq!(replay["error"]["code"], "consent_needs_user", "{replay}");
10478 pool.shutdown_all().await;
10479 }
10480
10481 #[tokio::test]
10482 async fn discover_all_rejects_combined_item_budget_without_publishing_any_family() {
10483 let sent = Arc::new(Mutex::new(Vec::new()));
10484 let tools = (0..MAX_MCP_CATALOG_ITEMS)
10485 .map(|i| serde_json::json!({"name": format!("tool_{i}"), "inputSchema": {}}))
10486 .collect::<Vec<_>>();
10487 let transport = ScriptedValueTransport {
10488 sent: Arc::clone(&sent),
10489 responses: VecDeque::from([
10490 json_frame(serde_json::json!({"jsonrpc":"2.0", "id":1, "result":{"tools":tools}})),
10491 json_frame(
10492 serde_json::json!({"jsonrpc":"2.0", "id":2, "result":{"resources":[{"uri":"file:///over-limit", "name":"over-limit"}]}}),
10493 ),
10494 ]),
10495 };
10496 let mut conn = test_connection(Box::new(transport));
10497 conn.server_capabilities = Some(McpServerCapabilities {
10498 tools: true,
10499 resources: true,
10500 prompts: true,
10501 });
10502 conn.tools = vec![
10503 serde_json::from_value(serde_json::json!({"name":"previous", "inputSchema":{}})).unwrap(),
10504 ];
10505 conn.resources = vec![
10506 serde_json::from_value(serde_json::json!({"uri":"file:///previous", "name":"previous"}))
10507 .unwrap(),
10508 ];
10509 conn.resource_templates = vec![
10510 serde_json::from_value(
10511 serde_json::json!({"uriTemplate":"file:///{key}", "name":"previous"}),
10512 )
10513 .unwrap(),
10514 ];
10515 conn.prompts = vec![serde_json::from_value(serde_json::json!({"name":"previous"})).unwrap()];
10516
10517 let error = conn
10518 .discover_all()
10519 .await
10520 .expect_err("combined count must refuse");
10521 assert!(
10522 error
10523 .to_string()
10524 .contains("resources/list exceeded the 4096-item")
10525 );
10526 assert_eq!(conn.tools.len(), 1);
10527 assert_eq!(conn.tools[0].name, "previous");
10528 assert_eq!(conn.resources[0].uri, "file:///previous");
10529 assert_eq!(conn.resource_templates[0].name, "previous");
10530 assert_eq!(conn.prompts[0].name, "previous");
10531 assert_eq!(
10532 sent.lock().unwrap().len(),
10533 2,
10534 "poisoned budget must prevent subsequent RPCs"
10535 );
10536 }
10537
10538 #[tokio::test]
10539 async fn discover_all_rejects_combined_page_budget_without_publishing() {
10540 let sent = Arc::new(Mutex::new(Vec::new()));
10541 let mut responses = VecDeque::new();
10542 for page in 1..=MAX_MCP_CATALOG_PAGES {
10543 let mut result = serde_json::json!({"tools":[]});
10544 if page < MAX_MCP_CATALOG_PAGES {
10545 result["nextCursor"] = serde_json::json!(format!("page_{page}"));
10546 }
10547 responses.push_back(json_frame(
10548 serde_json::json!({"jsonrpc":"2.0", "id":page, "result":result}),
10549 ));
10550 }
10551 responses.push_back(json_frame(serde_json::json!({"jsonrpc":"2.0", "id":MAX_MCP_CATALOG_PAGES+1, "result":{"resources":[]}})));
10552 let mut conn = test_connection(Box::new(ScriptedValueTransport {
10553 sent: Arc::clone(&sent),
10554 responses,
10555 }));
10556 conn.server_capabilities = Some(McpServerCapabilities {
10557 tools: true,
10558 resources: true,
10559 prompts: false,
10560 });
10561 conn.tools = vec![
10562 serde_json::from_value(serde_json::json!({"name":"previous", "inputSchema":{}})).unwrap(),
10563 ];
10564 let error = conn
10565 .discover_all()
10566 .await
10567 .expect_err("page budget must span families");
10568 assert!(
10569 error
10570 .to_string()
10571 .contains("resources/list exceeded the 64-page")
10572 );
10573 assert_eq!(conn.tools[0].name, "previous");
10574 assert_eq!(sent.lock().unwrap().len(), MAX_MCP_CATALOG_PAGES + 1);
10575 }
10576
10577 #[test]
10578 fn catalog_cursor_refusal_is_scoped_by_family_and_latched() {
10579 let mut budget = McpCatalogBudget::new();
10580 let page = serde_json::json!({"nextCursor":"shared-value"});
10581 budget.observe_page("tools/list", &page, 0).unwrap();
10582 budget.observe_page("resources/list", &page, 0).unwrap();
10583 let error = budget.observe_page("tools/list", &page, 0).unwrap_err();
10584 assert!(
10585 error
10586 .to_string()
10587 .contains("tools/list repeated pagination cursor")
10588 );
10589 let error = budget
10590 .observe_page("prompts/list", &serde_json::json!({}), 0)
10591 .unwrap_err();
10592 assert!(
10593 error
10594 .to_string()
10595 .contains("tools/list repeated pagination cursor")
10596 );
10597 }
10598
10599 #[test]
10600 fn catalog_byte_budget_cannot_reset_between_families() {
10601 let mut budget = McpCatalogBudget::new();
10602 // Put the real shared counter at the boundary without a 32MiB allocation.
10603 let page = serde_json::json!({});
10604 budget.bytes = MAX_MCP_CATALOG_BYTES - serde_json::to_vec(&page).unwrap().len();
10605 budget.observe_page("tools/list", &page, 0).unwrap();
10606 let error = budget.observe_page("resources/list", &page, 0).unwrap_err();
10607 assert!(
10608 error
10609 .to_string()
10610 .contains("resources/list exceeded the 33554432-byte aggregate")
10611 );
10612 assert!(budget.ensure_available().is_err());
10613 }
10614
10615 #[tokio::test]
10616 async fn discover_all_success_replaces_unavailable_old_families() {
10617 let mut conn = test_connection(Box::new(ScriptedValueTransport {
10618 sent: Arc::new(Mutex::new(Vec::new())),
10619 responses: VecDeque::from([json_frame(
10620 serde_json::json!({"jsonrpc":"2.0", "id":1, "result":{"tools":[]}}),
10621 )]),
10622 }));
10623 conn.server_capabilities = Some(McpServerCapabilities {
10624 tools: true,
10625 resources: false,
10626 prompts: false,
10627 });
10628 conn.resources = vec![
10629 serde_json::from_value(serde_json::json!({"uri":"file:///old", "name":"old"})).unwrap(),
10630 ];
10631 conn.prompts = vec![serde_json::from_value(serde_json::json!({"name":"old"})).unwrap()];
10632 conn.discover_all().await.unwrap();
10633 assert!(conn.tools.is_empty());
10634 assert!(conn.resources.is_empty());
10635 assert!(conn.prompts.is_empty());
10636 }
10637
10638 /// The selected Host and Rust default share the same real Rust OAuth store,
10639 /// refresh endpoint and connection recovery. Only loopback fixture credentials
10640 /// exist; no provider/browser or token bytes cross the extension RPC wire.
10641 #[tokio::test(flavor = "current_thread")]
10642 async fn http_backend_parity_reactive_refresh_reuses_exact_facade_id_once() {
10643 let _env = crate::test_support::lock_test_env();
10644 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
10645 let _loopback = lock_mcp_loopback_tests().await;
10646 let node = crate::extension_host::tests::node_for_tests("HTTP OAuth Host parity")
10647 .expect("Host parity requires Node");
10648 for backend in [McpBackend::Rust, McpBackend::Host] {
10649 let dir = tempfile::tempdir().unwrap();
10650 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
10651 let _policy = crate::plugins::activation::TestPolicyGuard::extension_host(true);
10652 let manager = Arc::new(crate::extension_host::ExtensionHostManager::new(
10653 crate::extension_host::ExtensionHostOptions {
10654 root: Some(dir.path().to_path_buf()),
10655 node_override: Some(node.clone()),
10656 ..Default::default()
10657 },
10658 ));
10659 let _manager = crate::extension_host::TestManagerGuard::install(Arc::clone(&manager));
10660 let mock = OAuthMcpMock::spawn().await;
10661 let url = mock.url();
10662 // The local expiry trusts this token; only the explicit peer 401 may
10663 // force refresh. A token-expiry fixture would not exercise this arm.
10664 seed_oauth_tokens(
10665 "wikiserver",
10666 &url,
10667 "cw-unaccepted-but-unexpired",
10668 "rt-fresh",
10669 Some(millis_from_now(3_600_000)),
10670 );
10671 let mut pool = McpPool::new(McpConfig {
10672 servers: [("wikiserver".into(), mock_oauth_server_config(mock.addr))].into(),
10673 ..Default::default()
10674 })
10675 .with_backend(backend);
10676 assert!(pool.connect_all().await.is_empty(), "{backend:?}");
10677 assert_eq!(mock.token_requests.load(AtomicOrdering::SeqCst), 1);
10678 let frames = mock.frames.lock().unwrap().clone();
10679 let initializes: Vec<_> = frames
10680 .iter()
10681 .filter(|frame| frame["method"] == "initialize")
10682 .collect();
10683 assert_eq!(initializes.len(), 2);
10684 assert_eq!(
10685 initializes[0], initializes[1],
10686 "same admitted frame, exact original Rust ID and params"
10687 );
10688 assert_eq!(initializes[0]["id"], "1");
10689 assert_eq!(
10690 frames
10691 .iter()
10692 .filter(|frame| frame["method"] == "notifications/initialized")
10693 .count(),
10694 1
10695 );
10696 pool.call_tool("mcp_wikiserver_wiki_lookup", json!({}))
10697 .await
10698 .unwrap();
10699 assert_eq!(
10700 mock.frames
10701 .lock()
10702 .unwrap()
10703 .iter()
10704 .filter(|frame| frame["method"] == "tools/call")
10705 .count(),
10706 1
10707 );
10708 pool.shutdown_all().await;
10709 manager.shutdown().await;
10710 mock.task.abort();
10711 }
10712 }
10713
10714 #[tokio::test(flavor = "current_thread")]
10715 async fn shared_http_refresh_revalidates_authority_before_any_second_write() {
10716 let _env = crate::test_support::lock_test_env();
10717 let dir = tempfile::tempdir().unwrap();
10718 let _home = crate::test_support::EnvVarGuard::set("CODEWHALE_HOME", dir.path());
10719 let _backend = crate::test_support::EnvVarGuard::set("CODEWHALE_SECRET_BACKEND", "file");
10720 let _loopback = lock_mcp_loopback_tests().await;
10721 let mock = OAuthMcpMock::spawn().await;
10722 let config = mock_oauth_server_config(mock.addr);
10723 let url = mock.url();
10724 seed_oauth_tokens(
10725 "wikiserver",
10726 &url,
10727 "cw-unaccepted-but-unexpired",
10728 "rt-fresh",
10729 Some(millis_from_now(3_600_000)),
10730 );
10731 let runtime = oauth::McpOAuthRuntime::from_server_config(
10732 "wikiserver",
10733 &config,
10734 reqwest::header::HeaderMap::new(),
10735 )
10736 .await
10737 .unwrap();
10738 let client = super::http_client::McpHttpClient::new(
10739 &url,
10740 false,
10741 false,
10742 false,
10743 None,
10744 Duration::from_secs(5),
10745 Duration::from_secs(5),
10746 )
10747 .unwrap()
10748 .with_mcp_auth(McpHttpAuth::from_config("wikiserver", &config, runtime));
10749 let checks = std::sync::atomic::AtomicUsize::new(0);
10750 let error = client
10751 .send_mcp_request(
10752 client.post(&url).body(
10753 json!({"jsonrpc":"2.0","id":"original-rust-id","method":"initialize","params":{}})
10754 .to_string(),
10755 ),
10756 true,
10757 false,
10758 true,
10759 || {
10760 // This is the actual post-refresh validation point, not a mock
10761 // success result: the token endpoint really answered first.
10762 if mock.token_requests.load(AtomicOrdering::SeqCst) != 0 {
10763 checks.fetch_add(1, AtomicOrdering::SeqCst);
10764 anyhow::bail!("fixture authority withdrawn during refresh");
10765 }
10766 Ok(())
10767 },
10768 |_| Ok(()),
10769 )
10770 .await
10771 .unwrap_err();
10772 assert!(error.to_string().contains("authority withdrawn"));
10773 assert_eq!(checks.load(AtomicOrdering::SeqCst), 1);
10774 assert_eq!(mock.token_requests.load(AtomicOrdering::SeqCst), 1);
10775 let frames = mock.frames.lock().unwrap();
10776 assert_eq!(
10777 frames.len(),
10778 1,
10779 "no second MCP network write after withdrawal"
10780 );
10781 assert_eq!(frames[0]["id"], "original-rust-id");
10782 drop(frames);
10783 mock.task.abort();
10784 }
10785
10786 #[test]
10787 fn configured_mcp_search_matches_real_names_and_bounds_the_admitted_set() {
10788 let mut config = McpConfig::default();
10789 for index in 0..12 {
10790 config.servers.insert(
10791 format!("server{index:02}"),
10792 serde_json::from_value(json!({"command":"unused"})).unwrap(),
10793 );
10794 }
10795 let pool = McpPool::new(config);
10796 assert!(
10797 pool.configured_servers_for_search(".*", "regex", |_| true)
10798 .unwrap()
10799 .is_empty()
10800 );
10801 assert!(
10802 pool.configured_servers_for_search("unrelated search", "bm25", |_| true)
10803 .unwrap()
10804 .is_empty()
10805 );
10806 let names = pool
10807 .configured_servers_for_search("mcp_.*", "regex", |name| name >= "server04")
10808 .unwrap();
10809 assert_eq!(names.len(), 8);
10810 assert_eq!(names[0], "server04");
10811 assert_eq!(names[7], "server11");
10812 assert_eq!(
10813 pool.configured_servers_for_search("^mcp_server05_actual_method$", "regex", |_| true)
10814 .unwrap(),
10815 ["server05"]
10816 );
10817 assert!(
10818 pool.configured_servers_for_search("mcp_[", "regex", |_| true)
10819 .is_err()
10820 );
10821 }
10822
10823 #[test]
10824 fn configured_mcp_search_excludes_disabled_and_namespace_denied_servers() {
10825 let mut config = McpConfig::default();
10826 config.servers.insert(
10827 "enabled".into(),
10828 serde_json::from_value(json!({"command":"unused"})).unwrap(),
10829 );
10830 config.servers.insert(
10831 "disabled".into(),
10832 serde_json::from_value(json!({"command":"unused", "enabled":false})).unwrap(),
10833 );
10834 config.servers.insert(
10835 "denied".into(),
10836 serde_json::from_value(json!({"command":"unused"})).unwrap(),
10837 );
10838 let pool = McpPool::new(config).with_disallowed_tools(vec!["mcp_denied_*".into()]);
10839 assert_eq!(
10840 pool.configured_servers_for_search("mcp_.*", "regex", |_| true)
10841 .unwrap(),
10842 ["enabled"]
10843 );
10844 }
10845
10845 lines RUST