use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex};
use tokio::sync::{mpsc, oneshot};
use tokio_util::sync::CancellationToken;
use mermaid_domain::{ApprovalChoice, ApprovalKind, Msg, ToolCallId, TurnId};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ApprovalDecision {
Approve,
ApproveAlways,
Deny,
}
impl From<ApprovalChoice> for ApprovalDecision {
fn from(choice: ApprovalChoice) -> Self {
match choice {
ApprovalChoice::Approve => Self::Approve,
ApprovalChoice::ApproveAlways => Self::ApproveAlways,
ApprovalChoice::Deny => Self::Deny,
}
}
}
struct PendingEntry {
tx: oneshot::Sender<ApprovalDecision>,
allowlist_key: String,
}
#[derive(Clone)]
pub struct ApprovalBroker {
pending: Arc<Mutex<HashMap<ToolCallId, PendingEntry>>>,
allowlist: Arc<Mutex<HashSet<String>>>,
msg_tx: mpsc::Sender<Msg>,
}
impl ApprovalBroker {
#[must_use]
pub fn new(msg_tx: mpsc::Sender<Msg>) -> Self {
Self {
pending: Arc::new(Mutex::new(HashMap::new())),
allowlist: Arc::new(Mutex::new(HashSet::new())),
msg_tx,
}
}
#[must_use]
pub fn is_allowlisted(&self, key: &str) -> bool {
self.allowlist
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.contains(key)
}
#[expect(clippy::too_many_arguments)]
pub async fn request(
&self,
token: &CancellationToken,
turn: TurnId,
call_id: ToolCallId,
tool: String,
risk: String,
kind: ApprovalKind,
prompt: String,
allowlist_key: String,
) -> ApprovalDecision {
let (tx, rx) = oneshot::channel();
self.pending
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.insert(
call_id,
PendingEntry {
tx,
allowlist_key: allowlist_key.clone(),
},
);
let sent = self
.msg_tx
.send(Msg::ApprovalRequested {
turn,
call_id,
tool,
risk,
kind,
prompt,
allowlist_scope: allowlist_key,
})
.await;
if sent.is_err() {
self.pending
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.remove(&call_id);
return ApprovalDecision::Deny;
}
tokio::select! {
biased;
_ = token.cancelled() => {
self.pending.lock().unwrap_or_else(|poisoned| poisoned.into_inner()).remove(&call_id);
ApprovalDecision::Deny
}
decision = rx => decision.unwrap_or(ApprovalDecision::Deny),
}
}
pub fn resolve(&self, call_id: ToolCallId, decision: ApprovalDecision) {
let entry = self
.pending
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.remove(&call_id);
if let Some(entry) = entry {
if decision == ApprovalDecision::ApproveAlways && !entry.allowlist_key.is_empty() {
self.allowlist
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.insert(entry.allowlist_key);
}
let _ = entry.tx.send(decision);
}
}
}
const FILE_TOOLS: &[&str] = &[
"read_file",
"write_file",
"edit_file",
"apply_patch",
"delete_file",
"create_directory",
];
const NON_ALLOWLISTABLE_TOOLS: &[&str] = &[
"type_text",
"press_key",
"click",
"mouse_move",
"scroll",
"mcp_proxy",
];
#[must_use]
pub fn extract_url_authority(raw: &str) -> Option<String> {
let input = raw.strip_prefix("web_fetch ").unwrap_or(raw).trim();
if input.is_empty() {
return None;
}
if let Ok(url) = reqwest::Url::parse(input)
&& let Some(host) = url.host_str()
{
let host_lower = host.to_ascii_lowercase();
if let Some(port) = url.port() {
return Some(format!("{host_lower}:{port}"));
}
return Some(host_lower);
}
if let Ok(url) = reqwest::Url::parse(&format!("https://{input}"))
&& let Some(host) = url.host_str()
{
let host_lower = host.to_ascii_lowercase();
if let Some(port) = url.port() {
return Some(format!("{host_lower}:{port}"));
}
return Some(host_lower);
}
None
}
#[must_use]
pub fn is_domain_allowed(allowed_domains: &[String], url_or_command: &str) -> bool {
if allowed_domains.is_empty() {
return false;
}
let Some(target_auth) = extract_url_authority(url_or_command) else {
return false;
};
allowed_domains.iter().any(|allowed| {
extract_url_authority(allowed)
.as_deref()
.is_some_and(|a| a.eq_ignore_ascii_case(&target_auth))
|| allowed.trim().eq_ignore_ascii_case(&target_auth)
})
}
#[must_use]
pub fn allowlist_key(tool: &str, command: Option<&str>, external_path: Option<&str>) -> String {
if NON_ALLOWLISTABLE_TOOLS.contains(&tool) {
return String::new();
}
if let Some(path) = external_path
&& FILE_TOOLS.contains(&tool)
{
let dir = std::path::Path::new(path)
.parent()
.map_or_else(|| path.to_string(), |d| d.display().to_string());
return format!("{tool}:{dir}");
}
if tool == "execute_command" {
let scope = external_path.map_or_else(String::new, |dir| format!("{dir}:"));
if let Some(cmd) = command {
let normalized = cmd.split_whitespace().collect::<Vec<_>>().join(" ");
if !normalized.is_empty() {
return format!("execute_command:{scope}{normalized}");
}
}
return format!("execute_command:{scope}")
.trim_end_matches(':')
.to_string();
}
if tool == "web_fetch" {
if let Some(cmd) = command
&& let Some(auth) = extract_url_authority(cmd)
{
return format!("web_fetch:{auth}");
}
return "web_fetch".to_string();
}
if tool == "web_search" {
return "web_search".to_string();
}
tool.to_string()
}
#[cfg(test)]
mod tests {
#[test]
fn execute_command_key_scopes_an_external_cwd() {
let inside = allowlist_key("execute_command", Some("make test"), None);
let a = allowlist_key("execute_command", Some("make test"), Some("/srv/a"));
let b = allowlist_key("execute_command", Some("make test"), Some("/srv/b"));
assert_eq!(inside, "execute_command:make test");
assert_eq!(a, "execute_command:/srv/a:make test");
assert_ne!(a, b);
assert_ne!(a, inside);
assert_eq!(
allowlist_key("execute_command", None, Some("/srv/a")),
"execute_command:/srv/a"
);
assert_eq!(
allowlist_key("execute_command", None, None),
"execute_command"
);
}
use super::*;
#[test]
fn allowlist_key_is_per_tool_with_full_command() {
assert_eq!(allowlist_key("write_file", None, None), "write_file");
assert_eq!(
allowlist_key("execute_command", Some("ls -la"), None),
"execute_command:ls -la"
);
assert_eq!(
allowlist_key("execute_command", None, None),
"execute_command"
);
}
#[test]
fn allowlist_key_distinguishes_argument_variants() {
assert_ne!(
allowlist_key("execute_command", Some("curl https://safe.example"), None),
allowlist_key("execute_command", Some("curl https://evil.example"), None),
);
assert_eq!(
allowlist_key("execute_command", Some("cargo build"), None),
allowlist_key("execute_command", Some("cargo build"), None),
);
assert_eq!(
allowlist_key("execute_command", Some("npm test"), None),
"execute_command:npm test"
);
assert_ne!(
allowlist_key("execute_command", Some("npm test"), None),
allowlist_key("execute_command", Some("npm run build"), None),
);
}
#[test]
fn web_fetch_and_search_allowlist_keys() {
assert_eq!(
allowlist_key(
"web_fetch",
Some("web_fetch https://docs.x.ai/docs/overview"),
None,
),
"web_fetch:docs.x.ai"
);
assert_eq!(
allowlist_key("web_fetch", Some("https://DOCS.X.AI/api"), None),
"web_fetch:docs.x.ai"
);
assert_eq!(
allowlist_key(
"web_fetch",
Some("web_fetch http://localhost:8080/query"),
None
),
"web_fetch:localhost:8080"
);
assert_eq!(allowlist_key("web_fetch", None, None), "web_fetch");
assert_eq!(
allowlist_key("web_search", Some("web_search rust documentation"), None),
"web_search"
);
assert_eq!(allowlist_key("web_search", None, None), "web_search");
}
#[test]
fn domain_allowlist_matching() {
let allowed = vec!["docs.x.ai".to_string(), "localhost:8080".to_string()];
assert!(is_domain_allowed(&allowed, "https://docs.x.ai/overview"));
assert!(is_domain_allowed(
&allowed,
"web_fetch https://DOCS.X.AI/api"
));
assert!(is_domain_allowed(&allowed, "http://localhost:8080/metrics"));
assert!(!is_domain_allowed(&allowed, "https://api.x.ai/v1"));
assert!(!is_domain_allowed(
&allowed,
"http://localhost:3000/metrics"
));
assert!(!is_domain_allowed(&[], "https://docs.x.ai/overview"));
}
#[test]
fn content_bearing_tools_are_non_allowlistable() {
for tool in [
"type_text",
"press_key",
"click",
"mouse_move",
"scroll",
"mcp_proxy",
] {
assert_eq!(
allowlist_key(tool, None, None),
"",
"{tool} must be non-allowlistable"
);
}
}
#[tokio::test]
async fn resolve_delivers_decision_and_approve_always_allowlists() {
let (tx, _rx) = mpsc::channel::<Msg>(8);
let broker = ApprovalBroker::new(tx);
let token = CancellationToken::new();
let b2 = broker.clone();
let handle = tokio::spawn(async move {
b2.request(
&CancellationToken::new(),
TurnId(1),
ToolCallId(1),
"execute_command".to_string(),
"shell_mutation".to_string(),
ApprovalKind::Shell,
"$ npm test".to_string(),
"execute_command:npm".to_string(),
)
.await
});
tokio::task::yield_now().await;
for _ in 0..100 {
broker.resolve(ToolCallId(1), ApprovalDecision::ApproveAlways);
if broker.is_allowlisted("execute_command:npm") {
break;
}
tokio::task::yield_now().await;
}
let decision = handle.await.unwrap();
assert_eq!(decision, ApprovalDecision::ApproveAlways);
assert!(broker.is_allowlisted("execute_command:npm"));
let _ = token; }
#[tokio::test]
async fn cancel_token_denies() {
let (tx, _rx) = mpsc::channel::<Msg>(8);
let broker = ApprovalBroker::new(tx);
let token = CancellationToken::new();
let token2 = token.clone();
let handle = tokio::spawn(async move {
broker
.request(
&token2,
TurnId(1),
ToolCallId(2),
"web_fetch".to_string(),
"network".to_string(),
ApprovalKind::Web,
"web_fetch https://x".to_string(),
"web_fetch".to_string(),
)
.await
});
tokio::task::yield_now().await;
token.cancel();
assert_eq!(handle.await.unwrap(), ApprovalDecision::Deny);
}
}