use std::collections::{HashMap, HashSet};
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::Mutex;
use leviath_mcp::{ToolDiscovery, ToolExecutor};
use leviath_providers::Tool;
use leviath_tools::{BuiltinTools, ToolContext};
use crate::config::{Config, ToolPolicy};
pub struct ToolRegistry {
pub builtins: Arc<BuiltinTools>,
pub mcp: Arc<Mutex<ToolExecutor>>,
pub mcp_tool_defs: Vec<Tool>,
pub builtin_names: HashSet<String>,
}
impl ToolRegistry {
pub async fn build(workdir: PathBuf, config: &Config) -> Self {
let ctx = ToolContext::new(workdir);
let builtins = Arc::new(BuiltinTools::new(ctx));
let builtin_names: HashSet<String> = builtins.names().into_iter().collect();
let mut mcp_executor = ToolExecutor::new();
let mut mcp_tool_defs: Vec<Tool> = Vec::new();
if !config.mcp_servers.is_empty() {
let mut discovery = ToolDiscovery::new();
let oauth = leviath_mcp::OAuthClient::new();
let store_path = leviath_mcp::AuthStore::default_path();
let now = unix_now_secs();
let credentials = credential_store_or_warn(crate::credentials::store_for(
config.security.credential_store,
));
for server_cfg in &config.mcp_servers {
let auth_header = match resolve_bearer(
&oauth,
&server_cfg.name,
store_path.as_deref(),
now,
credentials.as_deref(),
)
.await
{
Ok(header) => header,
Err(e) => {
tracing::warn!(server = %server_cfg.name, error = %e, "MCP auth unavailable - skipping");
continue;
}
};
let auth_was_resolved = auth_header.is_some();
match discovery
.discover_from_config_with_auth(
server_cfg,
auth_header,
&config.security.allow_env_vars,
)
.await
{
Ok((_tool_metas, mut client)) => {
if auth_was_resolved && let Some(path) = store_path.clone() {
client.set_refresher(std::sync::Arc::new(
leviath_mcp::StoredTokenRefresher::new(
server_cfg.name.clone(),
path,
),
));
}
let mut reserved: HashSet<String> = builtin_names.clone();
reserved.extend(mcp_tool_defs.iter().map(|t| t.name.clone()));
let advertised = mcp_executor.add_client_advertised(
server_cfg.name.clone(),
client,
&reserved,
);
for meta in advertised {
mcp_tool_defs.push(Tool {
name: meta.name,
description: meta.description,
parameters: meta.schema,
});
}
tracing::info!(server = %server_cfg.name, "Connected MCP server");
}
Err(e) => {
let span = tracing::warn_span!(
"mcp_server_connect_failed",
server = tracing::field::Empty,
error = tracing::field::Empty
);
let _enter = span.enter();
span.record("server", tracing::field::display(&server_cfg.name));
span.record("error", tracing::field::display(&e));
tracing::warn!("Failed to connect MCP server - skipping");
}
}
}
}
Self {
builtins,
mcp: Arc::new(Mutex::new(mcp_executor)),
mcp_tool_defs,
builtin_names,
}
}
pub fn all_tool_defs(&self) -> Vec<Tool> {
let mut tools = self.builtins.tool_defs();
tools.extend(BuiltinTools::subagent_tool_defs());
tools.extend_from_slice(&self.mcp_tool_defs);
tools
}
pub async fn shutdown(&self) {
let mut mcp = self.mcp.lock().await;
let _ = mcp.shutdown_all().await;
}
}
pub(crate) fn unix_now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
pub(crate) fn credential_store_or_warn(
resolved: crate::credentials::Resolved,
) -> Option<Box<dyn leviath_core::CredentialStore>> {
match resolved {
Ok(store) => store,
Err(e) => {
tracing::warn!("{e}. MCP servers needing OAuth will appear logged out.");
None
}
}
}
pub(crate) async fn resolve_bearer(
oauth: &leviath_mcp::OAuthClient,
server_name: &str,
store_path: Option<&std::path::Path>,
now: u64,
credentials: Option<&dyn leviath_core::CredentialStore>,
) -> anyhow::Result<Option<(String, String)>> {
match store_path {
Some(path) => {
oauth
.authorization_header_with(server_name, path, now, credentials)
.await
}
None => Ok(None),
}
}
pub fn default_tool_policy(tool_name: &str, is_builtin: bool) -> ToolPolicy {
match tool_name {
"read_file" | "list_dir" => ToolPolicy::Allow,
"write_file" | "edit_file" | "bash" => ToolPolicy::Ask,
"spawn_agent" | "check_agent" | "wait_for_agent" | "send_to_agent" | "kill_agent" => {
ToolPolicy::Allow
}
"ask_user_text" | "ask_user_choice" | "ask_user_confirm" | "edit_document" => {
ToolPolicy::Allow
}
_ => {
let _ = is_builtin;
ToolPolicy::Ask
}
}
}
fn restrictiveness(p: ToolPolicy) -> u8 {
match p {
ToolPolicy::Allow => 0,
ToolPolicy::Ask => 1,
ToolPolicy::Deny => 2,
}
}
fn stricter(a: ToolPolicy, b: ToolPolicy) -> ToolPolicy {
if restrictiveness(b) > restrictiveness(a) {
b
} else {
a
}
}
pub fn resolve_policy(
tool_name: &str,
is_builtin: bool,
launch_overrides: &HashMap<String, ToolPolicy>,
stage_permissions: &HashMap<String, String>,
agent_permissions: &HashMap<String, String>,
global_permissions: &HashMap<String, ToolPolicy>,
) -> ToolPolicy {
let ceiling = global_permissions.get(tool_name).copied();
let blueprint = stage_permissions
.get(tool_name)
.or_else(|| agent_permissions.get(tool_name))
.map(|s| parse_policy_str(s));
let configured = match (blueprint, ceiling) {
(Some(b), Some(c)) => stricter(b, c),
(Some(b), None) => b,
(None, Some(c)) => c,
(None, None) => default_tool_policy(tool_name, is_builtin),
};
if configured == ToolPolicy::Deny {
return ToolPolicy::Deny;
}
launch_overrides
.get(tool_name)
.or_else(|| launch_overrides.get("*"))
.copied()
.unwrap_or(configured)
}
pub fn session_approval_keys(tool_name: &str, arguments: &serde_json::Value) -> Vec<String> {
if leviath_tools::canonical_tool_name(tool_name) != "shell" {
return vec![tool_name.to_string()];
}
let Some(command) = arguments.get("command").and_then(|v| v.as_str()) else {
return Vec::new();
};
let segments = command_segments(command);
if segments.is_empty() {
return Vec::new();
}
let mut keys: Vec<String> = segments
.iter()
.filter_map(|seg| command_prefix(seg))
.map(|p| format!("shell:{p}"))
.collect();
keys.sort();
keys.dedup();
keys
}
fn command_segments(command: &str) -> Vec<String> {
let mut segments = Vec::new();
let mut current = String::new();
let mut rest = command;
while let Some((before, after_open)) = rest.split_once("$(") {
current.push_str(before);
let Some((inner, after)) = split_at_matching_paren(after_open) else {
return Vec::new(); };
segments.extend(command_segments(inner));
rest = after;
}
current.push_str(rest);
if current.contains('`') {
return Vec::new(); }
let without_redirects: String = current
.split(['>', '<'])
.next()
.unwrap_or_default()
.to_string();
segments.extend(
without_redirects
.split(['\n', ';', '&', '|'])
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_string),
);
segments
}
fn split_at_matching_paren(s: &str) -> Option<(&str, &str)> {
let mut depth = 0usize;
for (i, c) in s.char_indices() {
match c {
'(' => depth += 1,
')' if depth == 0 => return Some((s.split_at(i).0, s.split_at(i + 1).1)),
')' => depth -= 1,
_ => {}
}
}
None
}
fn command_prefix(command: &str) -> Option<String> {
let mut words = command.split_whitespace();
let program = words.next()?;
match words.next() {
Some(sub) if is_subcommand_like(sub) => Some(format!("{program} {sub}")),
_ => Some(program.to_string()),
}
}
fn is_subcommand_like(arg: &str) -> bool {
!arg.starts_with('-') && !arg.starts_with('"') && !arg.starts_with('\'') && !arg.contains('$')
}
fn parse_policy_str(s: &str) -> ToolPolicy {
match s.to_lowercase().as_str() {
"allow" => ToolPolicy::Allow,
"deny" => ToolPolicy::Deny,
_ => ToolPolicy::Ask,
}
}
#[cfg(test)]
mod mcp_registry_tests {
use super::*;
use crate::test_support::with_tracing;
use leviath_mcp::MCPServerConfig;
const STUB_INIT_AND_LIST: &str = r#"
import sys, json
def respond(id, result):
msg = json.dumps({"jsonrpc": "2.0", "id": id, "result": result})
sys.stdout.write(msg + "\n")
sys.stdout.flush()
for line in sys.stdin:
line = line.strip()
if not line:
continue
req = json.loads(line)
method = req.get("method", "")
id_ = req.get("id")
if method == "initialize":
respond(id_, {"capabilities": {"tools": {"listChanged": True}}, "protocolVersion": "2024-11-05"})
elif method == "notifications/initialized":
pass
elif method == "tools/list":
respond(id_, {"tools": [{"name": "echo", "description": "echo tool", "inputSchema": {}}]})
elif method == "tools/call":
args = req.get("params", {}).get("arguments", {})
if args.get("fail"):
respond(id_, {"content": [{"type": "text", "text": "it broke"}], "isError": True})
else:
respond(id_, {"content": [{"type": "text", "text": "echoed!"}], "isError": False})
else:
respond(id_, {"error": {"code": -32601, "message": "method not found"}})
"#;
fn config_with_mcp_server(command: &str, args: Vec<&str>) -> Config {
Config {
mcp_servers: vec![MCPServerConfig::stdio(
"stub-server",
command,
args.into_iter().map(String::from).collect(),
)],
..Config::default()
}
}
async fn with_temp_home<F, Fut, T>(body: F) -> T
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = T>,
{
let dir = tempfile::tempdir().unwrap();
temp_env::async_with_vars(
[("LEVIATH_HOME", Some(dir.path().to_str().unwrap()))],
body(),
)
.await
}
#[tokio::test]
async fn build_connects_mcp_server_and_registers_its_tools() {
with_tracing(|| {});
let registry = with_temp_home(|| async {
let config = config_with_mcp_server("python3", vec!["-c", STUB_INIT_AND_LIST]);
ToolRegistry::build(std::env::temp_dir(), &config).await
})
.await;
assert_eq!(registry.mcp_tool_defs.len(), 1);
assert_eq!(registry.mcp_tool_defs[0].name, "echo");
registry.shutdown().await;
}
#[tokio::test]
async fn build_advertises_two_servers_and_namespaces_a_collision() {
with_tracing(|| {});
let registry = with_temp_home(|| async {
let config = Config {
mcp_servers: vec![
MCPServerConfig::stdio(
"alpha",
"python3",
vec!["-c".to_string(), STUB_INIT_AND_LIST.to_string()],
),
MCPServerConfig::stdio(
"beta",
"python3",
vec!["-c".to_string(), STUB_INIT_AND_LIST.to_string()],
),
],
..Config::default()
};
ToolRegistry::build(std::env::temp_dir(), &config).await
})
.await;
let names: Vec<&str> = registry
.mcp_tool_defs
.iter()
.map(|t| t.name.as_str())
.collect();
assert!(names.contains(&"echo"), "names: {names:?}");
assert!(names.contains(&"beta__echo"), "names: {names:?}");
registry.shutdown().await;
}
async fn mock_http_mcp_server() -> String {
use axum::response::IntoResponse;
use axum::routing::post;
use axum::{Json, Router};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let app = Router::new().route(
"/mcp",
post(|body: String| async move {
let req: serde_json::Value = serde_json::from_str(&body).unwrap();
let id = req.get("id").cloned().unwrap_or(serde_json::json!(1));
let result = match req.get("method").and_then(|m| m.as_str()) {
Some("initialize") => {
serde_json::json!({"capabilities": {}, "protocolVersion": "2024-11-05"})
}
Some("tools/list") => {
serde_json::json!({"tools": [{"name": "remote_tool", "inputSchema": {}}]})
}
_ => serde_json::json!({}),
};
(
[(axum::http::header::CONTENT_TYPE, "application/json")],
Json(serde_json::json!({"jsonrpc": "2.0", "id": id, "result": result}))
.into_response()
.into_body(),
)
.into_response()
}),
);
tokio::spawn(std::future::IntoFuture::into_future(axum::serve(
listener, app,
)));
base
}
#[tokio::test]
async fn build_attaches_a_refresher_to_an_authenticated_http_server() {
with_tracing(|| {});
let base = mock_http_mcp_server().await;
let registry = with_temp_home(|| async {
let mut store = leviath_mcp::AuthStore::default();
store.set(
"remote",
leviath_mcp::ServerAuth {
access_token: "live-token".to_string(),
expires_at: u64::MAX,
..Default::default()
},
);
store
.save(&leviath_mcp::AuthStore::default_path().unwrap())
.unwrap();
let config = Config {
mcp_servers: vec![MCPServerConfig::http("remote", format!("{base}/mcp"))],
..Config::default()
};
ToolRegistry::build(std::env::temp_dir(), &config).await
})
.await;
assert_eq!(registry.mcp_tool_defs.len(), 1);
assert_eq!(registry.mcp_tool_defs[0].name, "remote_tool");
registry.shutdown().await;
}
#[tokio::test]
async fn build_skips_mcp_server_that_fails_to_connect() {
with_tracing(|| {});
let registry = with_temp_home(|| async {
let config = config_with_mcp_server("definitely-not-a-real-binary-xyz", vec![]);
ToolRegistry::build(std::env::temp_dir(), &config).await
})
.await;
assert!(registry.mcp_tool_defs.is_empty());
}
#[tokio::test]
async fn build_skips_http_server_whose_token_cannot_be_refreshed() {
with_tracing(|| {});
let registry = with_temp_home(|| async {
let mut store = leviath_mcp::AuthStore::default();
store.set(
"remote",
leviath_mcp::ServerAuth {
token_endpoint: "http://127.0.0.1:1/token".to_string(),
access_token: "expired".to_string(),
refresh_token: Some("good".to_string()),
expires_at: 1,
..Default::default()
},
);
store
.save(&leviath_mcp::AuthStore::default_path().unwrap())
.unwrap();
let config = Config {
mcp_servers: vec![MCPServerConfig::http("remote", "http://127.0.0.1:1/mcp")],
..Config::default()
};
ToolRegistry::build(std::env::temp_dir(), &config).await
})
.await;
assert!(registry.mcp_tool_defs.is_empty());
}
#[test]
fn an_unreachable_credential_store_warns_rather_than_failing_tool_setup() {
assert!(
credential_store_or_warn(Err("no keychain here".to_string())).is_none(),
"an unreachable store yields no credentials"
);
assert!(
credential_store_or_warn(Ok(None)).is_none(),
"and so does the file backend"
);
assert!(
credential_store_or_warn(Ok(Some(Box::new(leviath_core::MemoryStore::new()))))
.is_some()
);
}
#[tokio::test]
async fn resolve_bearer_without_a_store_is_none() {
let oauth = leviath_mcp::OAuthClient::new();
let header = resolve_bearer(&oauth, "srv", None, 0, None).await.unwrap();
assert!(header.is_none());
}
#[tokio::test]
async fn shutdown_with_no_servers_is_a_noop() {
let config = Config::default();
let registry = ToolRegistry::build(std::env::temp_dir(), &config).await;
registry.shutdown().await; }
}
#[cfg(test)]
mod policy_tests {
use super::*;
#[test]
fn test_default_policy_read_file() {
assert_eq!(default_tool_policy("read_file", true), ToolPolicy::Allow);
assert_eq!(default_tool_policy("list_dir", true), ToolPolicy::Allow);
}
#[test]
fn test_default_policy_write_tools() {
assert_eq!(default_tool_policy("write_file", true), ToolPolicy::Ask);
assert_eq!(default_tool_policy("edit_file", true), ToolPolicy::Ask);
assert_eq!(default_tool_policy("bash", true), ToolPolicy::Ask);
}
#[test]
fn test_default_policy_ask_user_tools_allow_by_default() {
assert_eq!(
default_tool_policy("ask_user_text", true),
ToolPolicy::Allow
);
assert_eq!(
default_tool_policy("ask_user_choice", true),
ToolPolicy::Allow
);
assert_eq!(
default_tool_policy("ask_user_confirm", true),
ToolPolicy::Allow
);
assert_eq!(
default_tool_policy("edit_document", true),
ToolPolicy::Allow
);
}
#[test]
fn test_resolve_policy_launch_override_wins() {
let mut launch = HashMap::new();
launch.insert("bash".to_string(), ToolPolicy::Allow);
let policy = resolve_policy(
"bash",
true,
&launch,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_yolo_wins() {
let mut launch = HashMap::new();
launch.insert("*".to_string(), ToolPolicy::Allow);
let policy = resolve_policy(
"bash",
true,
&launch,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_stage_may_tighten_global() {
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "deny".to_string());
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Allow);
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&stage,
&HashMap::new(),
&global,
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_stage_cannot_loosen_global() {
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "allow".to_string());
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Deny);
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&stage,
&HashMap::new(),
&global,
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_blueprint_free_when_user_silent() {
let mut agent = HashMap::new();
agent.insert("web_fetch".to_string(), "allow".to_string());
let policy = resolve_policy(
"web_fetch",
false,
&HashMap::new(),
&HashMap::new(),
&agent,
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_yolo_does_not_override_configured_deny() {
let mut launch = HashMap::new();
launch.insert("*".to_string(), ToolPolicy::Allow);
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Deny);
let policy = resolve_policy(
"bash",
true,
&launch,
&HashMap::new(),
&HashMap::new(),
&global,
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_named_allow_does_not_override_configured_deny() {
let mut launch = HashMap::new();
launch.insert("bash".to_string(), ToolPolicy::Allow);
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Deny);
let policy = resolve_policy(
"bash",
true,
&launch,
&HashMap::new(),
&HashMap::new(),
&global,
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_yolo_does_not_override_blueprint_deny() {
let mut launch = HashMap::new();
launch.insert("*".to_string(), ToolPolicy::Allow);
let mut agent = HashMap::new();
agent.insert("bash".to_string(), "deny".to_string());
let policy = resolve_policy(
"bash",
true,
&launch,
&HashMap::new(),
&agent,
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_yolo_still_collapses_ask_to_allow() {
let mut launch = HashMap::new();
launch.insert("*".to_string(), ToolPolicy::Allow);
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Ask);
let policy = resolve_policy(
"bash",
true,
&launch,
&HashMap::new(),
&HashMap::new(),
&global,
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_falls_through_to_default() {
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[test]
fn test_default_policy_unknown_tools() {
assert_eq!(default_tool_policy("unknown_tool", false), ToolPolicy::Ask);
assert_eq!(default_tool_policy("mcp_tool", false), ToolPolicy::Ask);
assert_eq!(default_tool_policy("custom_thing", true), ToolPolicy::Ask);
}
#[test]
fn test_resolve_policy_agent_cannot_loosen_global() {
let mut agent = HashMap::new();
agent.insert("bash".to_string(), "allow".to_string());
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Deny);
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&HashMap::new(),
&agent,
&global,
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_global_ask_bounds_blueprint_allow() {
let mut agent = HashMap::new();
agent.insert("write_file".to_string(), "allow".to_string());
let mut global = HashMap::new();
global.insert("write_file".to_string(), ToolPolicy::Ask);
let policy = resolve_policy(
"write_file",
true,
&HashMap::new(),
&HashMap::new(),
&agent,
&global,
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[test]
fn test_resolve_policy_launch_override_specific_beats_wildcard() {
let mut launch = HashMap::new();
launch.insert("bash".to_string(), ToolPolicy::Deny);
launch.insert("*".to_string(), ToolPolicy::Allow);
let policy = resolve_policy(
"bash",
true,
&launch,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_global_overrides_default() {
let mut global = HashMap::new();
global.insert("read_file".to_string(), ToolPolicy::Deny);
let policy = resolve_policy(
"read_file",
true,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&global,
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_stage_deny() {
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "deny".to_string());
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&stage,
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_stage_ask() {
let mut stage = HashMap::new();
stage.insert("read_file".to_string(), "ask".to_string());
let policy = resolve_policy(
"read_file",
true,
&HashMap::new(),
&stage,
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[test]
fn test_resolve_policy_unknown_stage_string_defaults_to_ask() {
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "unknown_policy".to_string());
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&stage,
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[test]
fn test_parse_policy_str_values() {
assert_eq!(parse_policy_str("allow"), ToolPolicy::Allow);
assert_eq!(parse_policy_str("Allow"), ToolPolicy::Allow);
assert_eq!(parse_policy_str("ALLOW"), ToolPolicy::Allow);
assert_eq!(parse_policy_str("deny"), ToolPolicy::Deny);
assert_eq!(parse_policy_str("Deny"), ToolPolicy::Deny);
assert_eq!(parse_policy_str("ask"), ToolPolicy::Ask);
assert_eq!(parse_policy_str("Ask"), ToolPolicy::Ask);
assert_eq!(parse_policy_str("anything_else"), ToolPolicy::Ask);
assert_eq!(parse_policy_str(""), ToolPolicy::Ask);
}
#[tokio::test]
async fn test_tool_registry_build_no_mcp() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
assert!(!registry.builtin_names.is_empty());
assert!(registry.mcp_tool_defs.is_empty());
}
#[tokio::test]
async fn test_tool_registry_all_tool_defs() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
let all_defs = registry.all_tool_defs();
assert!(!all_defs.is_empty());
let names: Vec<&str> = all_defs.iter().map(|t| t.name.as_str()).collect();
assert!(names.contains(&"read_file"));
}
#[tokio::test]
async fn test_tool_registry_builtin_names_consistent() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
let names_from_builtins: HashSet<String> = registry.builtins.names().into_iter().collect();
assert_eq!(registry.builtin_names, names_from_builtins);
}
fn shell_args(command: &str) -> serde_json::Value {
serde_json::json!({ "command": command })
}
fn keys(command: &str) -> Vec<String> {
session_approval_keys("shell", &shell_args(command))
}
#[test]
fn shell_approvals_are_keyed_on_the_command_prefix() {
assert_eq!(keys("ls -la"), ["shell:ls"]);
assert_eq!(keys("curl https://evil"), ["shell:curl https://evil"]);
assert_ne!(keys("ls -la"), keys("curl https://evil"));
}
#[test]
fn a_subcommand_is_part_of_the_prefix() {
assert_eq!(keys("git diff HEAD~1"), ["shell:git diff"]);
assert_ne!(keys("git diff HEAD~1"), keys("git push --force"));
}
#[test]
fn flags_are_not_part_of_the_prefix() {
assert_eq!(keys("cargo test --lib"), ["shell:cargo test"]);
assert_eq!(keys("cargo test --lib"), keys("cargo test --doc"));
assert_eq!(keys("ls -la"), keys("ls -l"));
}
#[test]
fn a_compound_line_grants_each_command_in_it() {
assert_eq!(keys("rm -rf __pycache__; ls -la"), ["shell:ls", "shell:rm"]);
assert_eq!(
keys(r#"test -f test.py && echo "created" || echo "missing""#),
["shell:echo", "shell:test"],
"quoted data must not split one program into two grants"
);
assert_eq!(
keys("python3 test.py | od -c | tail -5"),
["shell:od", "shell:python3 test.py", "shell:tail"]
);
}
#[test]
fn quoted_and_variable_arguments_are_not_part_of_the_key() {
assert_eq!(keys(r#"echo "exit code: $?""#), ["shell:echo"]);
assert_eq!(keys(r#"echo "done""#), keys(r#"echo "starting""#));
assert_eq!(keys("python3 test.py"), ["shell:python3 test.py"]);
assert_ne!(keys("python3 test.py"), keys("python3 evil.py"));
}
#[test]
fn approving_one_program_does_not_cover_a_line_that_runs_another() {
let granted: std::collections::HashSet<String> = keys("ls -la").into_iter().collect();
let attempted = keys("ls && curl https://evil");
assert!(
!attempted.iter().all(|k| granted.contains(k)),
"approving `ls` must not cover `ls && curl evil`: {attempted:?}"
);
assert!(attempted.iter().any(|k| k.starts_with("shell:curl")));
}
#[test]
fn a_substituted_command_gets_its_own_key() {
let k = keys("echo $(curl https://evil)");
assert!(k.iter().any(|k| k.starts_with("shell:curl")), "{k:?}");
assert!(k.iter().any(|k| k == "shell:echo"), "{k:?}");
let nested = keys("echo $(echo $(whoami))");
assert!(nested.iter().any(|k| k == "shell:whoami"), "{nested:?}");
}
#[test]
fn a_redirect_target_is_not_a_command() {
assert_eq!(
keys("cat /etc/passwd > /tmp/out"),
["shell:cat /etc/passwd"]
);
}
#[test]
fn an_unreadable_line_is_not_session_grantable() {
for command in [
"echo `whoami`", "echo $(unbalanced", " ", "&& ||", ] {
assert!(
keys(command).is_empty(),
"{command:?} must not be session-grantable"
);
}
}
#[test]
fn a_segment_with_no_program_has_no_prefix() {
assert_eq!(command_prefix(" "), None);
assert_eq!(command_prefix(""), None);
assert_eq!(command_prefix("ls"), Some("ls".to_string()));
}
#[test]
fn the_bash_alias_is_scoped_like_shell() {
assert_eq!(
session_approval_keys("bash", &shell_args("ls -la")),
["shell:ls"]
);
}
#[test]
fn other_tools_are_keyed_by_name() {
assert_eq!(
session_approval_keys("read_file", &serde_json::json!({ "path": "a" })),
["read_file"]
);
}
#[test]
fn a_shell_call_without_a_command_is_not_grantable() {
assert!(session_approval_keys("shell", &serde_json::json!({})).is_empty());
}
#[test]
fn test_resolve_policy_launch_overrides_stage_ask() {
let mut launch = HashMap::new();
launch.insert("bash".to_string(), ToolPolicy::Allow);
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "ask".to_string());
let policy = resolve_policy(
"bash",
true,
&launch,
&stage,
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_launch_cannot_override_stage_deny() {
let mut launch = HashMap::new();
launch.insert("bash".to_string(), ToolPolicy::Allow);
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "deny".to_string());
let policy = resolve_policy(
"bash",
true,
&launch,
&stage,
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_stage_overrides_agent() {
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "deny".to_string());
let mut agent = HashMap::new();
agent.insert("bash".to_string(), "allow".to_string());
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&stage,
&agent,
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_agent_overrides_global() {
let mut agent = HashMap::new();
agent.insert("write_file".to_string(), "deny".to_string());
let mut global = HashMap::new();
global.insert("write_file".to_string(), ToolPolicy::Allow);
let policy = resolve_policy(
"write_file",
true,
&HashMap::new(),
&HashMap::new(),
&agent,
&global,
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_wildcard_launch_with_missing_specific() {
let mut launch = HashMap::new();
launch.insert("*".to_string(), ToolPolicy::Allow);
let policy = resolve_policy(
"unknown_tool",
false,
&launch,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_mcp_tool_defaults_to_ask() {
let policy = resolve_policy(
"mcp_custom_tool",
false,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[test]
fn test_resolve_policy_read_file_default_is_allow() {
let policy = resolve_policy(
"read_file",
true,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_list_dir_default_is_allow() {
let policy = resolve_policy(
"list_dir",
true,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_write_file_default_is_ask() {
let policy = resolve_policy(
"write_file",
true,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[test]
fn test_resolve_policy_edit_file_default_is_ask() {
let policy = resolve_policy(
"edit_file",
true,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[tokio::test]
async fn test_tool_registry_shutdown_no_panic() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
registry.shutdown().await;
}
#[tokio::test]
async fn test_tool_registry_all_defs_includes_subagent() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
let all_defs = registry.all_tool_defs();
let names: Vec<&str> = all_defs.iter().map(|t| t.name.as_str()).collect();
assert!(names.contains(&"spawn_agent"));
}
#[test]
fn test_default_policy_search_is_ask() {
assert_eq!(default_tool_policy("search", true), ToolPolicy::Ask);
}
#[test]
fn test_default_policy_glob_is_ask() {
assert_eq!(default_tool_policy("glob", true), ToolPolicy::Ask);
}
#[test]
fn test_default_policy_http_request_is_ask() {
assert_eq!(default_tool_policy("http_request", true), ToolPolicy::Ask);
}
#[test]
fn test_default_policy_read_file_not_builtin_still_allow() {
assert_eq!(default_tool_policy("read_file", false), ToolPolicy::Allow);
}
#[test]
fn test_default_policy_list_dir_not_builtin_still_allow() {
assert_eq!(default_tool_policy("list_dir", false), ToolPolicy::Allow);
}
#[test]
fn test_resolve_policy_agent_deny() {
let mut agent = HashMap::new();
agent.insert("read_file".to_string(), "deny".to_string());
let policy = resolve_policy(
"read_file",
true,
&HashMap::new(),
&HashMap::new(),
&agent,
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_agent_unknown_string_defaults_to_ask() {
let mut agent = HashMap::new();
agent.insert("bash".to_string(), "foobar".to_string());
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&HashMap::new(),
&agent,
&HashMap::new(),
);
assert_eq!(policy, ToolPolicy::Ask);
}
#[test]
fn test_resolve_policy_global_allow_overrides_default_ask() {
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Allow);
let policy = resolve_policy(
"bash",
true,
&HashMap::new(),
&HashMap::new(),
&HashMap::new(),
&global,
);
assert_eq!(policy, ToolPolicy::Allow);
}
#[tokio::test]
async fn test_tool_registry_all_defs_includes_all_subagent_tools() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
let all_defs = registry.all_tool_defs();
let names: Vec<&str> = all_defs.iter().map(|t| t.name.as_str()).collect();
for expected in &[
"spawn_agent",
"check_agent",
"wait_for_agent",
"send_to_agent",
"kill_agent",
] {
assert!(names.contains(expected));
}
}
#[tokio::test]
async fn test_tool_registry_builtin_names_has_expected_tools() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
for name in &["read_file", "list_dir"] {
assert!(registry.builtin_names.contains(*name));
}
assert!(!registry.builtin_names.contains("spawn_agent"));
}
#[tokio::test]
async fn test_tool_registry_all_defs_no_mcp_when_none_configured() {
let config = Config::default();
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
assert!(registry.mcp_tool_defs.is_empty());
let all_defs = registry.all_tool_defs();
let builtin_count = registry.builtins.tool_defs().len();
let subagent_count = leviath_tools::BuiltinTools::subagent_tool_defs().len();
assert_eq!(all_defs.len(), builtin_count + subagent_count);
}
#[test]
fn test_resolve_policy_full_chain_deny_is_terminal() {
let mut launch = HashMap::new();
launch.insert("bash".to_string(), ToolPolicy::Allow);
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "deny".to_string());
let mut agent = HashMap::new();
agent.insert("bash".to_string(), "deny".to_string());
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Deny);
let policy = resolve_policy("bash", true, &launch, &stage, &agent, &global);
assert_eq!(policy, ToolPolicy::Deny);
}
#[test]
fn test_resolve_policy_full_chain_launch_relaxes_ask() {
let mut launch = HashMap::new();
launch.insert("bash".to_string(), ToolPolicy::Allow);
let mut stage = HashMap::new();
stage.insert("bash".to_string(), "ask".to_string());
let mut agent = HashMap::new();
agent.insert("bash".to_string(), "ask".to_string());
let mut global = HashMap::new();
global.insert("bash".to_string(), ToolPolicy::Ask);
let policy = resolve_policy("bash", true, &launch, &stage, &agent, &global);
assert_eq!(policy, ToolPolicy::Allow);
}
#[tokio::test]
async fn test_tool_registry_build_with_failing_mcp_server() {
use leviath_mcp::MCPServerConfig;
let bad_server = MCPServerConfig::stdio(
"bad-server",
"/nonexistent/binary/that/does/not/exist",
vec![],
);
let config = Config {
mcp_servers: vec![bad_server],
..Config::default()
};
let workdir = std::env::current_dir().unwrap();
let registry = ToolRegistry::build(workdir, &config).await;
assert!(registry.mcp_tool_defs.is_empty());
assert!(!registry.builtin_names.is_empty());
}
}