use std::collections::HashMap;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ToolPolicy {
#[default]
Auto,
OnDemand,
Always,
}
impl ToolPolicy {
pub fn max(self, other: Self) -> Self {
use ToolPolicy::*;
match (self, other) {
(Always, _) | (_, Always) => Always,
(OnDemand, _) | (_, OnDemand) => OnDemand,
_ => Auto,
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum ApprovalMode {
#[default]
Manual,
AllowList,
AutoRun,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ApprovalConfig {
#[serde(default)]
pub mode: ApprovalMode,
#[serde(default)]
pub allow_list: Vec<String>,
#[serde(default)]
pub tool_overrides: HashMap<String, ToolPolicy>,
}
pub const DEFAULT_TOOL_POLICIES: &[(&str, ToolPolicy)] = &[
("read", ToolPolicy::Auto),
("ls", ToolPolicy::Auto),
("grep", ToolPolicy::Auto),
("find", ToolPolicy::Auto),
("get_search_results", ToolPolicy::Auto),
("write", ToolPolicy::OnDemand),
("edit", ToolPolicy::OnDemand),
("exec", ToolPolicy::OnDemand),
("web_search", ToolPolicy::OnDemand),
("browser", ToolPolicy::OnDemand),
("mcp", ToolPolicy::OnDemand),
("a2a_delegate", ToolPolicy::OnDemand),
];
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn max_returns_stronger_policy() {
assert_eq!(
ToolPolicy::Auto.max(ToolPolicy::OnDemand),
ToolPolicy::OnDemand
);
assert_eq!(
ToolPolicy::OnDemand.max(ToolPolicy::Auto),
ToolPolicy::OnDemand
);
assert_eq!(ToolPolicy::Auto.max(ToolPolicy::Always), ToolPolicy::Always);
assert_eq!(ToolPolicy::Always.max(ToolPolicy::Auto), ToolPolicy::Always);
assert_eq!(
ToolPolicy::OnDemand.max(ToolPolicy::Always),
ToolPolicy::Always
);
assert_eq!(ToolPolicy::Auto.max(ToolPolicy::Auto), ToolPolicy::Auto);
}
#[test]
fn default_tool_policies_cover_core_tools() {
let names: Vec<_> = DEFAULT_TOOL_POLICIES.iter().map(|(n, _)| *n).collect();
for required in ["read", "write", "edit", "exec", "web_search", "grep", "ls"] {
assert!(
names.contains(&required),
"missing default policy for {required}"
);
}
}
#[test]
fn approval_config_defaults_to_manual_empty() {
let c = ApprovalConfig::default();
assert_eq!(c.mode, ApprovalMode::Manual);
assert!(c.allow_list.is_empty());
assert!(c.tool_overrides.is_empty());
}
#[test]
fn approval_mode_serde_kebab_case() {
let s = serde_json::to_string(&ApprovalMode::AllowList).unwrap();
assert_eq!(s, "\"allow-list\"");
let m: ApprovalMode = serde_json::from_str("\"auto-run\"").unwrap();
assert_eq!(m, ApprovalMode::AutoRun);
}
}