1use std::collections::{HashMap, HashSet};
4use std::path::PathBuf;
5use std::sync::Arc;
6use tokio::sync::Mutex;
7
8use leviath_mcp::{ToolDiscovery, ToolExecutor};
9use leviath_providers::Tool;
10use leviath_tools::{BuiltinTools, ToolContext};
11
12use crate::config::{Config, ToolPolicy};
13
14pub struct ToolRegistry {
19 pub builtins: Arc<BuiltinTools>,
20 pub mcp: Arc<Mutex<ToolExecutor>>,
21 pub mcp_tool_defs: Vec<Tool>,
22 pub builtin_names: HashSet<String>,
23}
24
25impl ToolRegistry {
26 pub async fn build(workdir: PathBuf, config: &Config) -> Self {
28 let ctx = ToolContext::new(workdir);
29 let builtins = Arc::new(BuiltinTools::new(ctx));
30 let builtin_names: HashSet<String> = builtins.names().into_iter().collect();
31
32 let mut mcp_executor = ToolExecutor::new();
33 let mut mcp_tool_defs: Vec<Tool> = Vec::new();
34
35 if !config.mcp_servers.is_empty() {
36 let mut discovery = ToolDiscovery::new();
37 let oauth = leviath_mcp::OAuthClient::new();
38 let store_path = leviath_mcp::AuthStore::default_path();
39 let now = unix_now_secs();
40 let credentials = credential_store_or_warn(crate::credentials::store_for(
46 config.security.credential_store,
47 ));
48 for server_cfg in &config.mcp_servers {
49 let auth_header = match resolve_bearer(
54 &oauth,
55 &server_cfg.name,
56 store_path.as_deref(),
57 now,
58 credentials.as_deref(),
59 )
60 .await
61 {
62 Ok(header) => header,
63 Err(e) => {
64 tracing::warn!(server = %server_cfg.name, error = %e, "MCP auth unavailable - skipping");
65 continue;
66 }
67 };
68 let auth_was_resolved = auth_header.is_some();
71 match discovery
72 .discover_from_config_with_auth(
73 server_cfg,
74 auth_header,
75 &config.security.allow_env_vars,
76 )
77 .await
78 {
79 Ok((_tool_metas, mut client)) => {
80 if auth_was_resolved && let Some(path) = store_path.clone() {
84 client.set_refresher(std::sync::Arc::new(
85 leviath_mcp::StoredTokenRefresher::new(
86 server_cfg.name.clone(),
87 path,
88 ),
89 ));
90 }
91 let mut reserved: HashSet<String> = builtin_names.clone();
96 reserved.extend(mcp_tool_defs.iter().map(|t| t.name.clone()));
97 let advertised = mcp_executor.add_client_advertised(
98 server_cfg.name.clone(),
99 client,
100 &reserved,
101 );
102 for meta in advertised {
103 mcp_tool_defs.push(Tool {
104 name: meta.name,
105 description: meta.description,
106 parameters: meta.schema,
107 });
108 }
109 tracing::info!(server = %server_cfg.name, "Connected MCP server");
110 }
111 Err(e) => {
112 let span = tracing::warn_span!(
113 "mcp_server_connect_failed",
114 server = tracing::field::Empty,
115 error = tracing::field::Empty
116 );
117 let _enter = span.enter();
118 span.record("server", tracing::field::display(&server_cfg.name));
119 span.record("error", tracing::field::display(&e));
120 tracing::warn!("Failed to connect MCP server - skipping");
121 }
122 }
123 }
124 }
125
126 Self {
127 builtins,
128 mcp: Arc::new(Mutex::new(mcp_executor)),
129 mcp_tool_defs,
130 builtin_names,
131 }
132 }
133
134 pub fn all_tool_defs(&self) -> Vec<Tool> {
136 let mut tools = self.builtins.tool_defs();
137 tools.extend(BuiltinTools::subagent_tool_defs());
138 tools.extend_from_slice(&self.mcp_tool_defs);
139 tools
140 }
141
142 pub async fn shutdown(&self) {
144 let mut mcp = self.mcp.lock().await;
145 let _ = mcp.shutdown_all().await;
150 }
151}
152
153pub(crate) fn unix_now_secs() -> u64 {
157 std::time::SystemTime::now()
158 .duration_since(std::time::UNIX_EPOCH)
159 .map(|d| d.as_secs())
160 .unwrap_or(0)
161}
162
163pub(crate) fn credential_store_or_warn(
176 resolved: crate::credentials::Resolved,
177) -> Option<Box<dyn leviath_core::CredentialStore>> {
178 match resolved {
179 Ok(store) => store,
180 Err(e) => {
181 tracing::warn!("{e}. MCP servers needing OAuth will appear logged out.");
182 None
183 }
184 }
185}
186
187pub(crate) async fn resolve_bearer(
188 oauth: &leviath_mcp::OAuthClient,
189 server_name: &str,
190 store_path: Option<&std::path::Path>,
191 now: u64,
192 credentials: Option<&dyn leviath_core::CredentialStore>,
193) -> anyhow::Result<Option<(String, String)>> {
194 match store_path {
195 Some(path) => {
196 oauth
197 .authorization_header_with(server_name, path, now, credentials)
198 .await
199 }
200 None => Ok(None),
201 }
202}
203
204pub fn default_tool_policy(tool_name: &str, is_builtin: bool) -> ToolPolicy {
208 match tool_name {
209 "read_file" | "list_dir" => ToolPolicy::Allow,
210 "write_file" | "edit_file" | "bash" => ToolPolicy::Ask,
211 "spawn_agent" | "check_agent" | "wait_for_agent" | "send_to_agent" | "kill_agent" => {
228 ToolPolicy::Allow
229 }
230 "ask_user_text" | "ask_user_choice" | "ask_user_confirm" | "edit_document" => {
234 ToolPolicy::Allow
235 }
236 _ => {
237 let _ = is_builtin;
239 ToolPolicy::Ask
240 }
241 }
242}
243
244fn restrictiveness(p: ToolPolicy) -> u8 {
246 match p {
247 ToolPolicy::Allow => 0,
248 ToolPolicy::Ask => 1,
249 ToolPolicy::Deny => 2,
250 }
251}
252
253fn stricter(a: ToolPolicy, b: ToolPolicy) -> ToolPolicy {
255 if restrictiveness(b) > restrictiveness(a) {
256 b
257 } else {
258 a
259 }
260}
261
262pub fn resolve_policy(
289 tool_name: &str,
290 is_builtin: bool,
291 launch_overrides: &HashMap<String, ToolPolicy>,
292 stage_permissions: &HashMap<String, String>,
293 agent_permissions: &HashMap<String, String>,
294 global_permissions: &HashMap<String, ToolPolicy>,
295) -> ToolPolicy {
296 let ceiling = global_permissions.get(tool_name).copied();
297
298 let blueprint = stage_permissions
300 .get(tool_name)
301 .or_else(|| agent_permissions.get(tool_name))
302 .map(|s| parse_policy_str(s));
303
304 let configured = match (blueprint, ceiling) {
305 (Some(b), Some(c)) => stricter(b, c),
306 (Some(b), None) => b,
307 (None, Some(c)) => c,
308 (None, None) => default_tool_policy(tool_name, is_builtin),
309 };
310
311 if configured == ToolPolicy::Deny {
313 return ToolPolicy::Deny;
314 }
315
316 launch_overrides
317 .get(tool_name)
318 .or_else(|| launch_overrides.get("*"))
319 .copied()
320 .unwrap_or(configured)
321}
322
323pub fn session_approval_keys(tool_name: &str, arguments: &serde_json::Value) -> Vec<String> {
349 if leviath_tools::canonical_tool_name(tool_name) != "shell" {
350 return vec![tool_name.to_string()];
351 }
352 let Some(command) = arguments.get("command").and_then(|v| v.as_str()) else {
353 return Vec::new();
354 };
355 let segments = command_segments(command);
356 if segments.is_empty() {
359 return Vec::new();
360 }
361 let mut keys: Vec<String> = segments
362 .iter()
363 .filter_map(|seg| command_prefix(seg))
364 .map(|p| format!("shell:{p}"))
365 .collect();
366 keys.sort();
367 keys.dedup();
368 keys
369}
370
371fn command_segments(command: &str) -> Vec<String> {
383 let mut segments = Vec::new();
388 let mut current = String::new();
389 let mut rest = command;
390
391 while let Some((before, after_open)) = rest.split_once("$(") {
395 current.push_str(before);
396 let Some((inner, after)) = split_at_matching_paren(after_open) else {
397 return Vec::new(); };
399 segments.extend(command_segments(inner));
401 rest = after;
402 }
403 current.push_str(rest);
404 if current.contains('`') {
405 return Vec::new(); }
407
408 let without_redirects: String = current
410 .split(['>', '<'])
411 .next()
412 .unwrap_or_default()
413 .to_string();
414 segments.extend(
415 without_redirects
416 .split(['\n', ';', '&', '|'])
417 .map(str::trim)
418 .filter(|s| !s.is_empty())
419 .map(str::to_string),
420 );
421 segments
422}
423
424fn split_at_matching_paren(s: &str) -> Option<(&str, &str)> {
432 let mut depth = 0usize;
433 for (i, c) in s.char_indices() {
434 match c {
435 '(' => depth += 1,
436 ')' if depth == 0 => return Some((s.split_at(i).0, s.split_at(i + 1).1)),
438 ')' => depth -= 1,
439 _ => {}
440 }
441 }
442 None
443}
444
445fn command_prefix(command: &str) -> Option<String> {
458 let mut words = command.split_whitespace();
459 let program = words.next()?;
460 match words.next() {
461 Some(sub) if is_subcommand_like(sub) => Some(format!("{program} {sub}")),
462 _ => Some(program.to_string()),
463 }
464}
465
466fn is_subcommand_like(arg: &str) -> bool {
469 !arg.starts_with('-') && !arg.starts_with('"') && !arg.starts_with('\'') && !arg.contains('$')
470}
471
472fn parse_policy_str(s: &str) -> ToolPolicy {
473 match s.to_lowercase().as_str() {
474 "allow" => ToolPolicy::Allow,
475 "deny" => ToolPolicy::Deny,
476 _ => ToolPolicy::Ask,
477 }
478}
479
480#[cfg(test)]
481mod mcp_registry_tests {
482 use super::*;
483 use crate::test_support::with_tracing;
484 use leviath_mcp::MCPServerConfig;
485
486 const STUB_INIT_AND_LIST: &str = r#"
492import sys, json
493
494def respond(id, result):
495 msg = json.dumps({"jsonrpc": "2.0", "id": id, "result": result})
496 sys.stdout.write(msg + "\n")
497 sys.stdout.flush()
498
499for line in sys.stdin:
500 line = line.strip()
501 if not line:
502 continue
503 req = json.loads(line)
504 method = req.get("method", "")
505 id_ = req.get("id")
506 if method == "initialize":
507 respond(id_, {"capabilities": {"tools": {"listChanged": True}}, "protocolVersion": "2024-11-05"})
508 elif method == "notifications/initialized":
509 pass
510 elif method == "tools/list":
511 respond(id_, {"tools": [{"name": "echo", "description": "echo tool", "inputSchema": {}}]})
512 elif method == "tools/call":
513 args = req.get("params", {}).get("arguments", {})
514 if args.get("fail"):
515 respond(id_, {"content": [{"type": "text", "text": "it broke"}], "isError": True})
516 else:
517 respond(id_, {"content": [{"type": "text", "text": "echoed!"}], "isError": False})
518 else:
519 respond(id_, {"error": {"code": -32601, "message": "method not found"}})
520"#;
521
522 fn config_with_mcp_server(command: &str, args: Vec<&str>) -> Config {
523 Config {
524 mcp_servers: vec![MCPServerConfig::stdio(
525 "stub-server",
526 command,
527 args.into_iter().map(String::from).collect(),
528 )],
529 ..Config::default()
530 }
531 }
532
533 async fn with_temp_home<F, Fut, T>(body: F) -> T
537 where
538 F: FnOnce() -> Fut,
539 Fut: std::future::Future<Output = T>,
540 {
541 let dir = tempfile::tempdir().unwrap();
542 temp_env::async_with_vars(
543 [("LEVIATH_HOME", Some(dir.path().to_str().unwrap()))],
544 body(),
545 )
546 .await
547 }
548
549 #[tokio::test]
550 async fn build_connects_mcp_server_and_registers_its_tools() {
551 with_tracing(|| {});
552 let registry = with_temp_home(|| async {
553 let config = config_with_mcp_server("python3", vec!["-c", STUB_INIT_AND_LIST]);
554 ToolRegistry::build(std::env::temp_dir(), &config).await
555 })
556 .await;
557
558 assert_eq!(registry.mcp_tool_defs.len(), 1);
559 assert_eq!(registry.mcp_tool_defs[0].name, "echo");
560
561 registry.shutdown().await;
562 }
563
564 #[tokio::test]
565 async fn build_advertises_two_servers_and_namespaces_a_collision() {
566 with_tracing(|| {});
571 let registry = with_temp_home(|| async {
572 let config = Config {
573 mcp_servers: vec![
574 MCPServerConfig::stdio(
575 "alpha",
576 "python3",
577 vec!["-c".to_string(), STUB_INIT_AND_LIST.to_string()],
578 ),
579 MCPServerConfig::stdio(
580 "beta",
581 "python3",
582 vec!["-c".to_string(), STUB_INIT_AND_LIST.to_string()],
583 ),
584 ],
585 ..Config::default()
586 };
587 ToolRegistry::build(std::env::temp_dir(), &config).await
588 })
589 .await;
590
591 let names: Vec<&str> = registry
592 .mcp_tool_defs
593 .iter()
594 .map(|t| t.name.as_str())
595 .collect();
596 assert!(names.contains(&"echo"), "names: {names:?}");
598 assert!(names.contains(&"beta__echo"), "names: {names:?}");
599 registry.shutdown().await;
600 }
601
602 async fn mock_http_mcp_server() -> String {
605 use axum::response::IntoResponse;
606 use axum::routing::post;
607 use axum::{Json, Router};
608 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
609 let base = format!("http://{}", listener.local_addr().unwrap());
610 let app = Router::new().route(
611 "/mcp",
612 post(|body: String| async move {
615 let req: serde_json::Value = serde_json::from_str(&body).unwrap();
616 let id = req.get("id").cloned().unwrap_or(serde_json::json!(1));
617 let result = match req.get("method").and_then(|m| m.as_str()) {
618 Some("initialize") => {
619 serde_json::json!({"capabilities": {}, "protocolVersion": "2024-11-05"})
620 }
621 Some("tools/list") => {
622 serde_json::json!({"tools": [{"name": "remote_tool", "inputSchema": {}}]})
623 }
624 _ => serde_json::json!({}),
625 };
626 (
627 [(axum::http::header::CONTENT_TYPE, "application/json")],
628 Json(serde_json::json!({"jsonrpc": "2.0", "id": id, "result": result}))
629 .into_response()
630 .into_body(),
631 )
632 .into_response()
633 }),
634 );
635 tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
636 listener, app,
637 )));
638 base
639 }
640
641 #[tokio::test]
642 async fn build_attaches_a_refresher_to_an_authenticated_http_server() {
643 with_tracing(|| {});
646 let base = mock_http_mcp_server().await;
647 let registry = with_temp_home(|| async {
648 let mut store = leviath_mcp::AuthStore::default();
650 store.set(
651 "remote",
652 leviath_mcp::ServerAuth {
653 access_token: "live-token".to_string(),
654 expires_at: u64::MAX,
655 ..Default::default()
656 },
657 );
658 store
659 .save(&leviath_mcp::AuthStore::default_path().unwrap())
660 .unwrap();
661
662 let config = Config {
663 mcp_servers: vec![MCPServerConfig::http("remote", format!("{base}/mcp"))],
664 ..Config::default()
665 };
666 ToolRegistry::build(std::env::temp_dir(), &config).await
667 })
668 .await;
669
670 assert_eq!(registry.mcp_tool_defs.len(), 1);
671 assert_eq!(registry.mcp_tool_defs[0].name, "remote_tool");
672 registry.shutdown().await;
673 }
674
675 #[tokio::test]
676 async fn build_skips_mcp_server_that_fails_to_connect() {
677 with_tracing(|| {});
681 let registry = with_temp_home(|| async {
682 let config = config_with_mcp_server("definitely-not-a-real-binary-xyz", vec![]);
683 ToolRegistry::build(std::env::temp_dir(), &config).await
684 })
685 .await;
686
687 assert!(registry.mcp_tool_defs.is_empty());
688 }
689
690 #[tokio::test]
691 async fn build_skips_http_server_whose_token_cannot_be_refreshed() {
692 with_tracing(|| {});
697 let registry = with_temp_home(|| async {
698 let mut store = leviath_mcp::AuthStore::default();
700 store.set(
701 "remote",
702 leviath_mcp::ServerAuth {
703 token_endpoint: "http://127.0.0.1:1/token".to_string(),
704 access_token: "expired".to_string(),
705 refresh_token: Some("good".to_string()),
706 expires_at: 1,
707 ..Default::default()
708 },
709 );
710 store
711 .save(&leviath_mcp::AuthStore::default_path().unwrap())
712 .unwrap();
713
714 let config = Config {
715 mcp_servers: vec![MCPServerConfig::http("remote", "http://127.0.0.1:1/mcp")],
716 ..Config::default()
717 };
718 ToolRegistry::build(std::env::temp_dir(), &config).await
719 })
720 .await;
721 assert!(registry.mcp_tool_defs.is_empty());
722 }
723
724 #[test]
727 fn an_unreachable_credential_store_warns_rather_than_failing_tool_setup() {
728 assert!(
729 credential_store_or_warn(Err("no keychain here".to_string())).is_none(),
730 "an unreachable store yields no credentials"
731 );
732 assert!(
733 credential_store_or_warn(Ok(None)).is_none(),
734 "and so does the file backend"
735 );
736 assert!(
737 credential_store_or_warn(Ok(Some(Box::new(leviath_core::MemoryStore::new()))))
738 .is_some()
739 );
740 }
741
742 #[tokio::test]
743 async fn resolve_bearer_without_a_store_is_none() {
744 let oauth = leviath_mcp::OAuthClient::new();
745 let header = resolve_bearer(&oauth, "srv", None, 0, None).await.unwrap();
746 assert!(header.is_none());
747 }
748
749 #[tokio::test]
750 async fn shutdown_with_no_servers_is_a_noop() {
751 let config = Config::default();
752 let registry = ToolRegistry::build(std::env::temp_dir(), &config).await;
753 registry.shutdown().await; }
755}
756
757#[cfg(test)]
758mod policy_tests {
759 use super::*;
760
761 #[test]
762 fn test_default_policy_read_file() {
763 assert_eq!(default_tool_policy("read_file", true), ToolPolicy::Allow);
764 assert_eq!(default_tool_policy("list_dir", true), ToolPolicy::Allow);
765 }
766
767 #[test]
768 fn test_default_policy_write_tools() {
769 assert_eq!(default_tool_policy("write_file", true), ToolPolicy::Ask);
770 assert_eq!(default_tool_policy("edit_file", true), ToolPolicy::Ask);
771 assert_eq!(default_tool_policy("bash", true), ToolPolicy::Ask);
772 }
773
774 #[test]
775 fn test_default_policy_ask_user_tools_allow_by_default() {
776 assert_eq!(
779 default_tool_policy("ask_user_text", true),
780 ToolPolicy::Allow
781 );
782 assert_eq!(
783 default_tool_policy("ask_user_choice", true),
784 ToolPolicy::Allow
785 );
786 assert_eq!(
787 default_tool_policy("ask_user_confirm", true),
788 ToolPolicy::Allow
789 );
790 assert_eq!(
791 default_tool_policy("edit_document", true),
792 ToolPolicy::Allow
793 );
794 }
795
796 #[test]
797 fn test_resolve_policy_launch_override_wins() {
798 let mut launch = HashMap::new();
799 launch.insert("bash".to_string(), ToolPolicy::Allow);
800 let policy = resolve_policy(
801 "bash",
802 true,
803 &launch,
804 &HashMap::new(),
805 &HashMap::new(),
806 &HashMap::new(),
807 );
808 assert_eq!(policy, ToolPolicy::Allow);
809 }
810
811 #[test]
812 fn test_resolve_policy_yolo_wins() {
813 let mut launch = HashMap::new();
814 launch.insert("*".to_string(), ToolPolicy::Allow);
815 let policy = resolve_policy(
816 "bash",
817 true,
818 &launch,
819 &HashMap::new(),
820 &HashMap::new(),
821 &HashMap::new(),
822 );
823 assert_eq!(policy, ToolPolicy::Allow);
824 }
825
826 #[test]
828 fn test_resolve_policy_stage_may_tighten_global() {
829 let mut stage = HashMap::new();
830 stage.insert("bash".to_string(), "deny".to_string());
831 let mut global = HashMap::new();
832 global.insert("bash".to_string(), ToolPolicy::Allow);
833 let policy = resolve_policy(
834 "bash",
835 true,
836 &HashMap::new(),
837 &stage,
838 &HashMap::new(),
839 &global,
840 );
841 assert_eq!(policy, ToolPolicy::Deny);
842 }
843
844 #[test]
850 fn test_resolve_policy_stage_cannot_loosen_global() {
851 let mut stage = HashMap::new();
852 stage.insert("bash".to_string(), "allow".to_string());
853 let mut global = HashMap::new();
854 global.insert("bash".to_string(), ToolPolicy::Deny);
855 let policy = resolve_policy(
856 "bash",
857 true,
858 &HashMap::new(),
859 &stage,
860 &HashMap::new(),
861 &global,
862 );
863 assert_eq!(policy, ToolPolicy::Deny);
864 }
865
866 #[test]
870 fn test_resolve_policy_blueprint_free_when_user_silent() {
871 let mut agent = HashMap::new();
872 agent.insert("web_fetch".to_string(), "allow".to_string());
873 let policy = resolve_policy(
874 "web_fetch",
875 false,
876 &HashMap::new(),
877 &HashMap::new(),
878 &agent,
879 &HashMap::new(),
880 );
881 assert_eq!(policy, ToolPolicy::Allow);
882 }
883
884 #[test]
889 fn test_yolo_does_not_override_configured_deny() {
890 let mut launch = HashMap::new();
891 launch.insert("*".to_string(), ToolPolicy::Allow);
892 let mut global = HashMap::new();
893 global.insert("bash".to_string(), ToolPolicy::Deny);
894 let policy = resolve_policy(
895 "bash",
896 true,
897 &launch,
898 &HashMap::new(),
899 &HashMap::new(),
900 &global,
901 );
902 assert_eq!(policy, ToolPolicy::Deny);
903 }
904
905 #[test]
907 fn test_named_allow_does_not_override_configured_deny() {
908 let mut launch = HashMap::new();
909 launch.insert("bash".to_string(), ToolPolicy::Allow);
910 let mut global = HashMap::new();
911 global.insert("bash".to_string(), ToolPolicy::Deny);
912 let policy = resolve_policy(
913 "bash",
914 true,
915 &launch,
916 &HashMap::new(),
917 &HashMap::new(),
918 &global,
919 );
920 assert_eq!(policy, ToolPolicy::Deny);
921 }
922
923 #[test]
926 fn test_yolo_does_not_override_blueprint_deny() {
927 let mut launch = HashMap::new();
928 launch.insert("*".to_string(), ToolPolicy::Allow);
929 let mut agent = HashMap::new();
930 agent.insert("bash".to_string(), "deny".to_string());
931 let policy = resolve_policy(
932 "bash",
933 true,
934 &launch,
935 &HashMap::new(),
936 &agent,
937 &HashMap::new(),
938 );
939 assert_eq!(policy, ToolPolicy::Deny);
940 }
941
942 #[test]
944 fn test_yolo_still_collapses_ask_to_allow() {
945 let mut launch = HashMap::new();
946 launch.insert("*".to_string(), ToolPolicy::Allow);
947 let mut global = HashMap::new();
948 global.insert("bash".to_string(), ToolPolicy::Ask);
949 let policy = resolve_policy(
950 "bash",
951 true,
952 &launch,
953 &HashMap::new(),
954 &HashMap::new(),
955 &global,
956 );
957 assert_eq!(policy, ToolPolicy::Allow);
958 }
959
960 #[test]
961 fn test_resolve_policy_falls_through_to_default() {
962 let policy = resolve_policy(
963 "bash",
964 true,
965 &HashMap::new(),
966 &HashMap::new(),
967 &HashMap::new(),
968 &HashMap::new(),
969 );
970 assert_eq!(policy, ToolPolicy::Ask);
971 }
972
973 #[test]
976 fn test_default_policy_unknown_tools() {
977 assert_eq!(default_tool_policy("unknown_tool", false), ToolPolicy::Ask);
978 assert_eq!(default_tool_policy("mcp_tool", false), ToolPolicy::Ask);
979 assert_eq!(default_tool_policy("custom_thing", true), ToolPolicy::Ask);
980 }
981
982 #[test]
986 fn test_resolve_policy_agent_cannot_loosen_global() {
987 let mut agent = HashMap::new();
988 agent.insert("bash".to_string(), "allow".to_string());
989 let mut global = HashMap::new();
990 global.insert("bash".to_string(), ToolPolicy::Deny);
991 let policy = resolve_policy(
992 "bash",
993 true,
994 &HashMap::new(),
995 &HashMap::new(),
996 &agent,
997 &global,
998 );
999 assert_eq!(policy, ToolPolicy::Deny);
1000 }
1001
1002 #[test]
1005 fn test_resolve_policy_global_ask_bounds_blueprint_allow() {
1006 let mut agent = HashMap::new();
1007 agent.insert("write_file".to_string(), "allow".to_string());
1008 let mut global = HashMap::new();
1009 global.insert("write_file".to_string(), ToolPolicy::Ask);
1010 let policy = resolve_policy(
1011 "write_file",
1012 true,
1013 &HashMap::new(),
1014 &HashMap::new(),
1015 &agent,
1016 &global,
1017 );
1018 assert_eq!(policy, ToolPolicy::Ask);
1019 }
1020
1021 #[test]
1022 fn test_resolve_policy_launch_override_specific_beats_wildcard() {
1023 let mut launch = HashMap::new();
1024 launch.insert("bash".to_string(), ToolPolicy::Deny);
1025 launch.insert("*".to_string(), ToolPolicy::Allow);
1026 let policy = resolve_policy(
1027 "bash",
1028 true,
1029 &launch,
1030 &HashMap::new(),
1031 &HashMap::new(),
1032 &HashMap::new(),
1033 );
1034 assert_eq!(policy, ToolPolicy::Deny);
1036 }
1037
1038 #[test]
1039 fn test_resolve_policy_global_overrides_default() {
1040 let mut global = HashMap::new();
1041 global.insert("read_file".to_string(), ToolPolicy::Deny);
1042 let policy = resolve_policy(
1043 "read_file",
1044 true,
1045 &HashMap::new(),
1046 &HashMap::new(),
1047 &HashMap::new(),
1048 &global,
1049 );
1050 assert_eq!(policy, ToolPolicy::Deny);
1051 }
1052
1053 #[test]
1054 fn test_resolve_policy_stage_deny() {
1055 let mut stage = HashMap::new();
1056 stage.insert("bash".to_string(), "deny".to_string());
1057 let policy = resolve_policy(
1058 "bash",
1059 true,
1060 &HashMap::new(),
1061 &stage,
1062 &HashMap::new(),
1063 &HashMap::new(),
1064 );
1065 assert_eq!(policy, ToolPolicy::Deny);
1066 }
1067
1068 #[test]
1069 fn test_resolve_policy_stage_ask() {
1070 let mut stage = HashMap::new();
1071 stage.insert("read_file".to_string(), "ask".to_string());
1072 let policy = resolve_policy(
1073 "read_file",
1074 true,
1075 &HashMap::new(),
1076 &stage,
1077 &HashMap::new(),
1078 &HashMap::new(),
1079 );
1080 assert_eq!(policy, ToolPolicy::Ask);
1081 }
1082
1083 #[test]
1084 fn test_resolve_policy_unknown_stage_string_defaults_to_ask() {
1085 let mut stage = HashMap::new();
1086 stage.insert("bash".to_string(), "unknown_policy".to_string());
1087 let policy = resolve_policy(
1088 "bash",
1089 true,
1090 &HashMap::new(),
1091 &stage,
1092 &HashMap::new(),
1093 &HashMap::new(),
1094 );
1095 assert_eq!(policy, ToolPolicy::Ask);
1096 }
1097
1098 #[test]
1101 fn test_parse_policy_str_values() {
1102 assert_eq!(parse_policy_str("allow"), ToolPolicy::Allow);
1103 assert_eq!(parse_policy_str("Allow"), ToolPolicy::Allow);
1104 assert_eq!(parse_policy_str("ALLOW"), ToolPolicy::Allow);
1105 assert_eq!(parse_policy_str("deny"), ToolPolicy::Deny);
1106 assert_eq!(parse_policy_str("Deny"), ToolPolicy::Deny);
1107 assert_eq!(parse_policy_str("ask"), ToolPolicy::Ask);
1108 assert_eq!(parse_policy_str("Ask"), ToolPolicy::Ask);
1109 assert_eq!(parse_policy_str("anything_else"), ToolPolicy::Ask);
1110 assert_eq!(parse_policy_str(""), ToolPolicy::Ask);
1111 }
1112
1113 #[tokio::test]
1116 async fn test_tool_registry_build_no_mcp() {
1117 let config = Config::default();
1118 let workdir = std::env::current_dir().unwrap();
1119 let registry = ToolRegistry::build(workdir, &config).await;
1120
1121 assert!(!registry.builtin_names.is_empty());
1123 assert!(registry.mcp_tool_defs.is_empty());
1125 }
1126
1127 #[tokio::test]
1128 async fn test_tool_registry_all_tool_defs() {
1129 let config = Config::default();
1130 let workdir = std::env::current_dir().unwrap();
1131 let registry = ToolRegistry::build(workdir, &config).await;
1132
1133 let all_defs = registry.all_tool_defs();
1134 assert!(!all_defs.is_empty());
1135
1136 let names: Vec<&str> = all_defs.iter().map(|t| t.name.as_str()).collect();
1138 assert!(names.contains(&"read_file"));
1139 }
1140
1141 #[tokio::test]
1142 async fn test_tool_registry_builtin_names_consistent() {
1143 let config = Config::default();
1144 let workdir = std::env::current_dir().unwrap();
1145 let registry = ToolRegistry::build(workdir, &config).await;
1146
1147 let names_from_builtins: HashSet<String> = registry.builtins.names().into_iter().collect();
1149 assert_eq!(registry.builtin_names, names_from_builtins);
1150 }
1151
1152 fn shell_args(command: &str) -> serde_json::Value {
1157 serde_json::json!({ "command": command })
1158 }
1159
1160 fn keys(command: &str) -> Vec<String> {
1161 session_approval_keys("shell", &shell_args(command))
1162 }
1163
1164 #[test]
1167 fn shell_approvals_are_keyed_on_the_command_prefix() {
1168 assert_eq!(keys("ls -la"), ["shell:ls"]);
1169 assert_eq!(keys("curl https://evil"), ["shell:curl https://evil"]);
1170 assert_ne!(keys("ls -la"), keys("curl https://evil"));
1171 }
1172
1173 #[test]
1176 fn a_subcommand_is_part_of_the_prefix() {
1177 assert_eq!(keys("git diff HEAD~1"), ["shell:git diff"]);
1178 assert_ne!(keys("git diff HEAD~1"), keys("git push --force"));
1179 }
1180
1181 #[test]
1184 fn flags_are_not_part_of_the_prefix() {
1185 assert_eq!(keys("cargo test --lib"), ["shell:cargo test"]);
1186 assert_eq!(keys("cargo test --lib"), keys("cargo test --doc"));
1187 assert_eq!(keys("ls -la"), keys("ls -l"));
1188 }
1189
1190 #[test]
1195 fn a_compound_line_grants_each_command_in_it() {
1196 assert_eq!(keys("rm -rf __pycache__; ls -la"), ["shell:ls", "shell:rm"]);
1197 assert_eq!(
1198 keys(r#"test -f test.py && echo "created" || echo "missing""#),
1199 ["shell:echo", "shell:test"],
1200 "quoted data must not split one program into two grants"
1201 );
1202 assert_eq!(
1203 keys("python3 test.py | od -c | tail -5"),
1204 ["shell:od", "shell:python3 test.py", "shell:tail"]
1205 );
1206 }
1207
1208 #[test]
1212 fn quoted_and_variable_arguments_are_not_part_of_the_key() {
1213 assert_eq!(keys(r#"echo "exit code: $?""#), ["shell:echo"]);
1214 assert_eq!(keys(r#"echo "done""#), keys(r#"echo "starting""#));
1215 assert_eq!(keys("python3 test.py"), ["shell:python3 test.py"]);
1218 assert_ne!(keys("python3 test.py"), keys("python3 evil.py"));
1219 }
1220
1221 #[test]
1225 fn approving_one_program_does_not_cover_a_line_that_runs_another() {
1226 let granted: std::collections::HashSet<String> = keys("ls -la").into_iter().collect();
1227 let attempted = keys("ls && curl https://evil");
1228 assert!(
1229 !attempted.iter().all(|k| granted.contains(k)),
1230 "approving `ls` must not cover `ls && curl evil`: {attempted:?}"
1231 );
1232 assert!(attempted.iter().any(|k| k.starts_with("shell:curl")));
1234 }
1235
1236 #[test]
1240 fn a_substituted_command_gets_its_own_key() {
1241 let k = keys("echo $(curl https://evil)");
1242 assert!(k.iter().any(|k| k.starts_with("shell:curl")), "{k:?}");
1243 assert!(k.iter().any(|k| k == "shell:echo"), "{k:?}");
1244 let nested = keys("echo $(echo $(whoami))");
1246 assert!(nested.iter().any(|k| k == "shell:whoami"), "{nested:?}");
1247 }
1248
1249 #[test]
1252 fn a_redirect_target_is_not_a_command() {
1253 assert_eq!(
1254 keys("cat /etc/passwd > /tmp/out"),
1255 ["shell:cat /etc/passwd"]
1256 );
1257 }
1258
1259 #[test]
1263 fn an_unreadable_line_is_not_session_grantable() {
1264 for command in [
1265 "echo `whoami`", "echo $(unbalanced", " ", "&& ||", ] {
1270 assert!(
1271 keys(command).is_empty(),
1272 "{command:?} must not be session-grantable"
1273 );
1274 }
1275 }
1276
1277 #[test]
1280 fn a_segment_with_no_program_has_no_prefix() {
1281 assert_eq!(command_prefix(" "), None);
1282 assert_eq!(command_prefix(""), None);
1283 assert_eq!(command_prefix("ls"), Some("ls".to_string()));
1284 }
1285
1286 #[test]
1289 fn the_bash_alias_is_scoped_like_shell() {
1290 assert_eq!(
1291 session_approval_keys("bash", &shell_args("ls -la")),
1292 ["shell:ls"]
1293 );
1294 }
1295
1296 #[test]
1299 fn other_tools_are_keyed_by_name() {
1300 assert_eq!(
1301 session_approval_keys("read_file", &serde_json::json!({ "path": "a" })),
1302 ["read_file"]
1303 );
1304 }
1305
1306 #[test]
1309 fn a_shell_call_without_a_command_is_not_grantable() {
1310 assert!(session_approval_keys("shell", &serde_json::json!({})).is_empty());
1311 }
1312
1313 #[test]
1315 fn test_resolve_policy_launch_overrides_stage_ask() {
1316 let mut launch = HashMap::new();
1317 launch.insert("bash".to_string(), ToolPolicy::Allow);
1318 let mut stage = HashMap::new();
1319 stage.insert("bash".to_string(), "ask".to_string());
1320 let policy = resolve_policy(
1321 "bash",
1322 true,
1323 &launch,
1324 &stage,
1325 &HashMap::new(),
1326 &HashMap::new(),
1327 );
1328 assert_eq!(policy, ToolPolicy::Allow);
1329 }
1330
1331 #[test]
1334 fn test_resolve_policy_launch_cannot_override_stage_deny() {
1335 let mut launch = HashMap::new();
1336 launch.insert("bash".to_string(), ToolPolicy::Allow);
1337 let mut stage = HashMap::new();
1338 stage.insert("bash".to_string(), "deny".to_string());
1339 let policy = resolve_policy(
1340 "bash",
1341 true,
1342 &launch,
1343 &stage,
1344 &HashMap::new(),
1345 &HashMap::new(),
1346 );
1347 assert_eq!(policy, ToolPolicy::Deny);
1348 }
1349
1350 #[test]
1351 fn test_resolve_policy_stage_overrides_agent() {
1352 let mut stage = HashMap::new();
1353 stage.insert("bash".to_string(), "deny".to_string());
1354 let mut agent = HashMap::new();
1355 agent.insert("bash".to_string(), "allow".to_string());
1356 let policy = resolve_policy(
1357 "bash",
1358 true,
1359 &HashMap::new(),
1360 &stage,
1361 &agent,
1362 &HashMap::new(),
1363 );
1364 assert_eq!(policy, ToolPolicy::Deny);
1365 }
1366
1367 #[test]
1368 fn test_resolve_policy_agent_overrides_global() {
1369 let mut agent = HashMap::new();
1370 agent.insert("write_file".to_string(), "deny".to_string());
1371 let mut global = HashMap::new();
1372 global.insert("write_file".to_string(), ToolPolicy::Allow);
1373 let policy = resolve_policy(
1374 "write_file",
1375 true,
1376 &HashMap::new(),
1377 &HashMap::new(),
1378 &agent,
1379 &global,
1380 );
1381 assert_eq!(policy, ToolPolicy::Deny);
1382 }
1383
1384 #[test]
1385 fn test_resolve_policy_wildcard_launch_with_missing_specific() {
1386 let mut launch = HashMap::new();
1387 launch.insert("*".to_string(), ToolPolicy::Allow);
1388 let policy = resolve_policy(
1390 "unknown_tool",
1391 false,
1392 &launch,
1393 &HashMap::new(),
1394 &HashMap::new(),
1395 &HashMap::new(),
1396 );
1397 assert_eq!(policy, ToolPolicy::Allow);
1398 }
1399
1400 #[test]
1401 fn test_resolve_policy_mcp_tool_defaults_to_ask() {
1402 let policy = resolve_policy(
1403 "mcp_custom_tool",
1404 false,
1405 &HashMap::new(),
1406 &HashMap::new(),
1407 &HashMap::new(),
1408 &HashMap::new(),
1409 );
1410 assert_eq!(policy, ToolPolicy::Ask);
1411 }
1412
1413 #[test]
1414 fn test_resolve_policy_read_file_default_is_allow() {
1415 let policy = resolve_policy(
1416 "read_file",
1417 true,
1418 &HashMap::new(),
1419 &HashMap::new(),
1420 &HashMap::new(),
1421 &HashMap::new(),
1422 );
1423 assert_eq!(policy, ToolPolicy::Allow);
1424 }
1425
1426 #[test]
1427 fn test_resolve_policy_list_dir_default_is_allow() {
1428 let policy = resolve_policy(
1429 "list_dir",
1430 true,
1431 &HashMap::new(),
1432 &HashMap::new(),
1433 &HashMap::new(),
1434 &HashMap::new(),
1435 );
1436 assert_eq!(policy, ToolPolicy::Allow);
1437 }
1438
1439 #[test]
1440 fn test_resolve_policy_write_file_default_is_ask() {
1441 let policy = resolve_policy(
1442 "write_file",
1443 true,
1444 &HashMap::new(),
1445 &HashMap::new(),
1446 &HashMap::new(),
1447 &HashMap::new(),
1448 );
1449 assert_eq!(policy, ToolPolicy::Ask);
1450 }
1451
1452 #[test]
1453 fn test_resolve_policy_edit_file_default_is_ask() {
1454 let policy = resolve_policy(
1455 "edit_file",
1456 true,
1457 &HashMap::new(),
1458 &HashMap::new(),
1459 &HashMap::new(),
1460 &HashMap::new(),
1461 );
1462 assert_eq!(policy, ToolPolicy::Ask);
1463 }
1464
1465 #[tokio::test]
1466 async fn test_tool_registry_shutdown_no_panic() {
1467 let config = Config::default();
1468 let workdir = std::env::current_dir().unwrap();
1469 let registry = ToolRegistry::build(workdir, &config).await;
1470 registry.shutdown().await;
1471 }
1472
1473 #[tokio::test]
1474 async fn test_tool_registry_all_defs_includes_subagent() {
1475 let config = Config::default();
1476 let workdir = std::env::current_dir().unwrap();
1477 let registry = ToolRegistry::build(workdir, &config).await;
1478 let all_defs = registry.all_tool_defs();
1479 let names: Vec<&str> = all_defs.iter().map(|t| t.name.as_str()).collect();
1480 assert!(names.contains(&"spawn_agent"));
1482 }
1483
1484 #[test]
1487 fn test_default_policy_search_is_ask() {
1488 assert_eq!(default_tool_policy("search", true), ToolPolicy::Ask);
1489 }
1490
1491 #[test]
1492 fn test_default_policy_glob_is_ask() {
1493 assert_eq!(default_tool_policy("glob", true), ToolPolicy::Ask);
1494 }
1495
1496 #[test]
1497 fn test_default_policy_http_request_is_ask() {
1498 assert_eq!(default_tool_policy("http_request", true), ToolPolicy::Ask);
1499 }
1500
1501 #[test]
1502 fn test_default_policy_read_file_not_builtin_still_allow() {
1503 assert_eq!(default_tool_policy("read_file", false), ToolPolicy::Allow);
1505 }
1506
1507 #[test]
1508 fn test_default_policy_list_dir_not_builtin_still_allow() {
1509 assert_eq!(default_tool_policy("list_dir", false), ToolPolicy::Allow);
1510 }
1511
1512 #[test]
1515 fn test_resolve_policy_agent_deny() {
1516 let mut agent = HashMap::new();
1517 agent.insert("read_file".to_string(), "deny".to_string());
1518 let policy = resolve_policy(
1519 "read_file",
1520 true,
1521 &HashMap::new(),
1522 &HashMap::new(),
1523 &agent,
1524 &HashMap::new(),
1525 );
1526 assert_eq!(policy, ToolPolicy::Deny);
1527 }
1528
1529 #[test]
1532 fn test_resolve_policy_agent_unknown_string_defaults_to_ask() {
1533 let mut agent = HashMap::new();
1534 agent.insert("bash".to_string(), "foobar".to_string());
1535 let policy = resolve_policy(
1536 "bash",
1537 true,
1538 &HashMap::new(),
1539 &HashMap::new(),
1540 &agent,
1541 &HashMap::new(),
1542 );
1543 assert_eq!(policy, ToolPolicy::Ask);
1544 }
1545
1546 #[test]
1549 fn test_resolve_policy_global_allow_overrides_default_ask() {
1550 let mut global = HashMap::new();
1551 global.insert("bash".to_string(), ToolPolicy::Allow);
1552 let policy = resolve_policy(
1553 "bash",
1554 true,
1555 &HashMap::new(),
1556 &HashMap::new(),
1557 &HashMap::new(),
1558 &global,
1559 );
1560 assert_eq!(policy, ToolPolicy::Allow);
1561 }
1562
1563 #[tokio::test]
1564 async fn test_tool_registry_all_defs_includes_all_subagent_tools() {
1565 let config = Config::default();
1566 let workdir = std::env::current_dir().unwrap();
1567 let registry = ToolRegistry::build(workdir, &config).await;
1568 let all_defs = registry.all_tool_defs();
1569 let names: Vec<&str> = all_defs.iter().map(|t| t.name.as_str()).collect();
1570
1571 for expected in &[
1572 "spawn_agent",
1573 "check_agent",
1574 "wait_for_agent",
1575 "send_to_agent",
1576 "kill_agent",
1577 ] {
1578 assert!(names.contains(expected));
1579 }
1580 }
1581
1582 #[tokio::test]
1585 async fn test_tool_registry_builtin_names_has_expected_tools() {
1586 let config = Config::default();
1587 let workdir = std::env::current_dir().unwrap();
1588 let registry = ToolRegistry::build(workdir, &config).await;
1589
1590 for name in &["read_file", "list_dir"] {
1592 assert!(registry.builtin_names.contains(*name));
1593 }
1594
1595 assert!(!registry.builtin_names.contains("spawn_agent"));
1597 }
1598
1599 #[tokio::test]
1602 async fn test_tool_registry_all_defs_no_mcp_when_none_configured() {
1603 let config = Config::default();
1604 let workdir = std::env::current_dir().unwrap();
1605 let registry = ToolRegistry::build(workdir, &config).await;
1606 assert!(registry.mcp_tool_defs.is_empty());
1607
1608 let all_defs = registry.all_tool_defs();
1610 let builtin_count = registry.builtins.tool_defs().len();
1611 let subagent_count = leviath_tools::BuiltinTools::subagent_tool_defs().len();
1612 assert_eq!(all_defs.len(), builtin_count + subagent_count);
1613 }
1614
1615 #[test]
1621 fn test_resolve_policy_full_chain_deny_is_terminal() {
1622 let mut launch = HashMap::new();
1623 launch.insert("bash".to_string(), ToolPolicy::Allow);
1624 let mut stage = HashMap::new();
1625 stage.insert("bash".to_string(), "deny".to_string());
1626 let mut agent = HashMap::new();
1627 agent.insert("bash".to_string(), "deny".to_string());
1628 let mut global = HashMap::new();
1629 global.insert("bash".to_string(), ToolPolicy::Deny);
1630
1631 let policy = resolve_policy("bash", true, &launch, &stage, &agent, &global);
1632 assert_eq!(policy, ToolPolicy::Deny);
1633 }
1634
1635 #[test]
1638 fn test_resolve_policy_full_chain_launch_relaxes_ask() {
1639 let mut launch = HashMap::new();
1640 launch.insert("bash".to_string(), ToolPolicy::Allow);
1641 let mut stage = HashMap::new();
1642 stage.insert("bash".to_string(), "ask".to_string());
1643 let mut agent = HashMap::new();
1644 agent.insert("bash".to_string(), "ask".to_string());
1645 let mut global = HashMap::new();
1646 global.insert("bash".to_string(), ToolPolicy::Ask);
1647
1648 let policy = resolve_policy("bash", true, &launch, &stage, &agent, &global);
1649 assert_eq!(policy, ToolPolicy::Allow);
1650 }
1651
1652 #[tokio::test]
1656 async fn test_tool_registry_build_with_failing_mcp_server() {
1657 use leviath_mcp::MCPServerConfig;
1658
1659 let bad_server = MCPServerConfig::stdio(
1660 "bad-server",
1661 "/nonexistent/binary/that/does/not/exist",
1662 vec![],
1663 );
1664 let config = Config {
1665 mcp_servers: vec![bad_server],
1666 ..Config::default()
1667 };
1668
1669 let workdir = std::env::current_dir().unwrap();
1670 let registry = ToolRegistry::build(workdir, &config).await;
1672
1673 assert!(registry.mcp_tool_defs.is_empty());
1675 assert!(!registry.builtin_names.is_empty());
1677 }
1678
1679 }