use std::path::PathBuf;
use crate::domain::{ApprovalKind, ToolOutcome};
use crate::providers::{ApprovalBroker, ApprovalDecision, allowlist_key};
use crate::runtime::{
ActionRequest, NewApproval, PolicyDecision, PolicyEngine, RiskClass, RuntimeStore,
create_checkpoint_for_task, run_plugin_hooks,
};
use super::super::ctx::ExecContext;
pub enum Gate {
Proceed { risk: RiskClass },
Block(ToolOutcome),
}
pub async fn gate_external(
ctx: &ExecContext,
tool: &'static str,
category: crate::runtime::ToolCategory,
summary: String,
args: &serde_json::Value,
) -> Option<ToolOutcome> {
let request = ActionRequest::new(tool, category, summary);
let pending = serde_json::json!({ "tool": tool, "args": args });
match gate(ctx, request, &[], pending, false).await {
Gate::Block(outcome) => Some(outcome),
Gate::Proceed { .. } => None,
}
}
pub async fn gate(
ctx: &ExecContext,
request: ActionRequest,
checkpoint_paths: &[PathBuf],
pending_action: serde_json::Value,
replayable: bool,
) -> Gate {
let decision = PolicyEngine::new(ctx.safety_mode)
.with_overrides(ctx.config.safety.overrides.clone())
.decide(&request);
match decision {
PolicyDecision::Allow { risk, .. } => Gate::Proceed { risk },
PolicyDecision::Ask { risk, checkpoint } => {
if let Some(broker) = &ctx.approval {
inline_decision(ctx, broker, &request, risk, None).await
} else if !replayable {
tracing::debug!(
tool = %request.tool,
"policy Ask on non-replayable tool with no approval UI; proceeding",
);
Gate::Proceed { risk }
} else {
block_for_approval(
ctx,
&request,
checkpoint,
checkpoint_paths,
pending_action,
risk,
None,
)
}
},
PolicyDecision::Classify { risk, checkpoint } => {
let verdict = match &ctx.classifier {
Some(classifier) => {
let vreq = crate::providers::VetRequest {
tool: request.tool.clone(),
summary: request.summary.clone(),
command: request.command.clone(),
path: request.path.clone(),
intent: ctx.intent.clone(),
workdir: ctx.workdir.display().to_string(),
turn: ctx.turn,
token: ctx.token.clone(),
};
classifier.vet(&vreq).await
},
None => crate::providers::VetVerdict::escalate("no Auto-mode classifier available"),
};
if verdict.allow {
Gate::Proceed { risk }
} else if let Some(broker) = &ctx.approval {
inline_decision(ctx, broker, &request, risk, Some(verdict.reason)).await
} else if replayable {
block_for_approval(
ctx,
&request,
checkpoint,
checkpoint_paths,
pending_action,
risk,
Some(verdict.reason),
)
} else {
Gate::Block(ToolOutcome::error(
format!(
"{} blocked by Auto-mode safety review: {}",
request.summary, verdict.reason
),
0.0,
))
}
},
PolicyDecision::Deny { reason, .. } => Gate::Block(ToolOutcome::error(
format!("{} blocked by policy: {}", request.summary, reason),
0.0,
)),
}
}
async fn inline_decision(
ctx: &ExecContext,
broker: &ApprovalBroker,
request: &ActionRequest,
risk: RiskClass,
classifier_reason: Option<String>,
) -> Gate {
let key = allowlist_key(&request.tool, request.command.as_deref());
if broker.is_allowlisted(&key) {
return Gate::Proceed { risk };
}
let kind = if classifier_reason.is_some() {
ApprovalKind::Classify
} else {
approval_kind(request.category)
};
let prompt = format_approval_body(request, classifier_reason.as_deref());
let decision = broker
.request(
&ctx.token,
ctx.turn,
ctx.call_id,
request.tool.clone(),
risk.as_str().to_string(),
kind,
prompt,
key,
)
.await;
match decision {
ApprovalDecision::Approve | ApprovalDecision::ApproveAlways => Gate::Proceed { risk },
ApprovalDecision::Deny => Gate::Block(ToolOutcome::error(
format!("{} — denied by you", request.summary),
0.0,
)),
}
}
fn format_approval_body(request: &ActionRequest, classifier_reason: Option<&str>) -> String {
let mut body = if let Some(cmd) = &request.command {
format!("$ {}", cmd)
} else if let Some(path) = &request.path {
format!("{} ({})", path, request.summary)
} else {
request.summary.clone()
};
if let Some(reason) = classifier_reason {
body.push_str(&format!("\n\nAuto-review flagged this: {}", reason));
}
body
}
fn approval_kind(category: crate::runtime::ToolCategory) -> ApprovalKind {
use crate::runtime::ToolCategory as C;
match category {
C::Edit => ApprovalKind::FileMutation,
C::Shell | C::Git | C::Process => ApprovalKind::Shell,
C::Web | C::Network | C::ExternalDirectory => ApprovalKind::Web,
C::Mcp => ApprovalKind::Mcp,
C::Subagent => ApprovalKind::Subagent,
C::ComputerUse => ApprovalKind::ComputerUse,
C::Read | C::Memory => ApprovalKind::Shell,
}
}
#[allow(clippy::too_many_arguments)]
fn block_for_approval(
ctx: &ExecContext,
request: &ActionRequest,
checkpoint: bool,
checkpoint_paths: &[PathBuf],
pending_action: serde_json::Value,
risk: RiskClass,
classifier_reason: Option<String>,
) -> Gate {
let checkpoint_id = if checkpoint && ctx.config.safety.checkpoint_on_mutation {
match create_checkpoint_for_task(
&ctx.workdir,
checkpoint_paths,
Some(pending_action.clone()),
ctx.task_id.clone(),
) {
Ok(manifest) => Some(manifest.id),
Err(error) => {
return Gate::Block(ToolOutcome::error(
format!(
"{} checkpoint failed before approval: {}",
request.summary, error
),
0.0,
));
},
}
} else {
None
};
let args_summary = request
.command
.clone()
.or_else(|| request.path.clone())
.unwrap_or_else(|| request.summary.clone());
let pending_action_json = serde_json::to_string(&pending_action).ok();
let tool = request.tool.clone();
let risk_str = risk.as_str().to_string();
let proposed_action = match &classifier_reason {
Some(reason) => format!("{} [auto-review: {}]", request.summary, reason),
None => request.summary.clone(),
};
let approval_id = RuntimeStore::open_default()
.and_then(|store| {
let approval = store.approvals().create(NewApproval {
task_id: ctx.task_id.clone(),
proposed_action: proposed_action.clone(),
risk_classification: risk_str.clone(),
policy_decision: "ask".to_string(),
args_summary: Some(args_summary),
checkpoint_id: checkpoint_id.clone(),
pending_action_json,
})?;
if let Some(checkpoint_id) = checkpoint_id.as_deref() {
let _ = store
.checkpoints()
.set_approval(checkpoint_id, &approval.id);
}
let _ = run_plugin_hooks(
"approval_requested",
&serde_json::json!({
"id": approval.id.clone(),
"task_id": approval.task_id.clone(),
"tool": tool,
"risk": risk_str,
"checkpoint_id": checkpoint_id.clone(),
}),
);
Ok(approval)
})
.map(|approval| approval.id)
.ok();
Gate::Block(ToolOutcome::error(
format!(
"Approval required for {}{}",
request.summary,
approval_id
.map(|id| format!(" (approval {})", id))
.unwrap_or_default()
),
0.0,
))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::domain::{ToolCallId, TurnId};
use crate::providers::ctx::ProgressEvent;
use crate::runtime::{SafetyMode, ToolCategory};
use std::path::PathBuf;
use std::sync::Arc;
fn ctx(mode: SafetyMode) -> ExecContext {
let mut config = crate::app::Config::default();
config.safety.mode = mode;
let (tx, _rx) = tokio::sync::mpsc::channel::<ProgressEvent>(4);
ExecContext::new(
tokio_util::sync::CancellationToken::new(),
tx,
ToolCallId(1),
TurnId(1),
PathBuf::from("."),
Arc::new(config),
String::new(),
None,
mode,
None,
None,
None,
)
}
struct StubClassifier {
allow: bool,
}
#[async_trait::async_trait]
impl crate::providers::AutoClassifier for StubClassifier {
async fn vet(&self, _req: &crate::providers::VetRequest) -> crate::providers::VetVerdict {
if self.allow {
crate::providers::VetVerdict::allow()
} else {
crate::providers::VetVerdict::escalate("stub: misaligned")
}
}
}
fn ctx_auto(classifier: Option<Arc<dyn crate::providers::AutoClassifier>>) -> ExecContext {
let mut config = crate::app::Config::default();
config.safety.mode = SafetyMode::Auto;
let (tx, _rx) = tokio::sync::mpsc::channel::<ProgressEvent>(4);
ExecContext::new(
tokio_util::sync::CancellationToken::new(),
tx,
ToolCallId(1),
TurnId(1),
PathBuf::from("."),
Arc::new(config),
String::new(),
None,
SafetyMode::Auto,
Some("fetch the changelog".to_string()),
classifier,
None,
)
}
#[tokio::test]
async fn readonly_blocks_external_tools() {
let ctx = ctx(SafetyMode::ReadOnly);
for (tool, cat) in [
("web_fetch", ToolCategory::Web),
("mcp_proxy", ToolCategory::Mcp),
("agent", ToolCategory::Subagent),
("click", ToolCategory::ComputerUse),
("memory", ToolCategory::Memory),
] {
assert!(
gate_external(&ctx, tool, cat, tool.to_string(), &serde_json::json!({}))
.await
.is_some(),
"ReadOnly must block {tool}",
);
}
}
#[tokio::test]
async fn memory_writes_ungated_except_readonly() {
for mode in [SafetyMode::Ask, SafetyMode::Auto, SafetyMode::FullAccess] {
let ctx = ctx(mode);
assert!(
gate_external(
&ctx,
"memory",
ToolCategory::Memory,
"memory remember".to_string(),
&serde_json::json!({"action": "remember"}),
)
.await
.is_none(),
"memory must proceed without approval in {mode:?}",
);
}
let ctx = ctx(SafetyMode::ReadOnly);
assert!(
gate_external(
&ctx,
"memory",
ToolCategory::Memory,
"memory remember".to_string(),
&serde_json::json!({"action": "remember"}),
)
.await
.is_some(),
"read-only must block memory writes",
);
}
#[tokio::test]
async fn full_access_allows_external_tools() {
let ctx = ctx(SafetyMode::FullAccess);
assert!(
gate_external(
&ctx,
"web_fetch",
ToolCategory::Web,
"web_fetch".to_string(),
&serde_json::json!({}),
)
.await
.is_none()
);
}
#[tokio::test]
async fn auto_classifier_allow_proceeds() {
let ctx = ctx_auto(Some(Arc::new(StubClassifier { allow: true })));
assert!(
gate_external(
&ctx,
"web_fetch",
ToolCategory::Web,
"web_fetch".to_string(),
&serde_json::json!({}),
)
.await
.is_none(),
"ALLOW verdict should let the action proceed",
);
}
#[tokio::test]
async fn auto_classifier_escalate_blocks() {
let ctx = ctx_auto(Some(Arc::new(StubClassifier { allow: false })));
assert!(
gate_external(
&ctx,
"web_fetch",
ToolCategory::Web,
"web_fetch".to_string(),
&serde_json::json!({}),
)
.await
.is_some(),
"ESCALATE verdict should block a non-replayable tool",
);
}
#[tokio::test]
async fn auto_without_classifier_fails_safe() {
let ctx = ctx_auto(None);
assert!(
gate_external(
&ctx,
"web_fetch",
ToolCategory::Web,
"web_fetch".to_string(),
&serde_json::json!({}),
)
.await
.is_some(),
"missing classifier must fail safe (block), not allow",
);
}
fn ctx_with_broker(broker: crate::providers::ApprovalBroker) -> ExecContext {
let mut config = crate::app::Config::default();
config.safety.mode = SafetyMode::Ask;
let (tx, _rx) = tokio::sync::mpsc::channel::<ProgressEvent>(4);
ExecContext::new(
tokio_util::sync::CancellationToken::new(),
tx,
ToolCallId(7),
TurnId(1),
PathBuf::from("."),
Arc::new(config),
String::new(),
None,
SafetyMode::Ask,
None,
None,
Some(broker),
)
}
fn shell_request(cmd: &str) -> ActionRequest {
let mut req = ActionRequest::new("execute_command", ToolCategory::Shell, cmd);
req.command = Some(cmd.to_string());
req
}
#[tokio::test]
async fn inline_ask_approve_proceeds() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<crate::domain::Msg>(8);
let broker = crate::providers::ApprovalBroker::new(tx);
let ctx = ctx_with_broker(broker.clone());
let handle = tokio::spawn(async move {
gate(
&ctx,
shell_request("npm test"),
&[],
serde_json::json!({}),
true,
)
.await
});
let call_id = match rx.recv().await.expect("approval requested") {
crate::domain::Msg::ApprovalRequested { call_id, .. } => call_id,
other => panic!("expected ApprovalRequested, got {other:?}"),
};
broker.resolve(call_id, crate::providers::ApprovalDecision::Approve);
assert!(matches!(handle.await.unwrap(), Gate::Proceed { .. }));
}
#[tokio::test]
async fn inline_ask_deny_blocks() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<crate::domain::Msg>(8);
let broker = crate::providers::ApprovalBroker::new(tx);
let ctx = ctx_with_broker(broker.clone());
let handle = tokio::spawn(async move {
gate(
&ctx,
shell_request("rm -rf node_modules"),
&[],
serde_json::json!({}),
true,
)
.await
});
let call_id = match rx.recv().await.expect("approval requested") {
crate::domain::Msg::ApprovalRequested { call_id, .. } => call_id,
other => panic!("expected ApprovalRequested, got {other:?}"),
};
broker.resolve(call_id, crate::providers::ApprovalDecision::Deny);
assert!(matches!(handle.await.unwrap(), Gate::Block(_)));
}
#[tokio::test]
async fn inline_allowlisted_skips_prompt() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<crate::domain::Msg>(8);
let broker = crate::providers::ApprovalBroker::new(tx);
let ctx1 = ctx_with_broker(broker.clone());
let b1 = broker.clone();
let h1 = tokio::spawn(async move {
gate(
&ctx1,
shell_request("npm run build"),
&[],
serde_json::json!({}),
true,
)
.await
});
let id = match rx.recv().await.expect("first prompt") {
crate::domain::Msg::ApprovalRequested { call_id, .. } => call_id,
other => panic!("got {other:?}"),
};
b1.resolve(id, crate::providers::ApprovalDecision::ApproveAlways);
assert!(matches!(h1.await.unwrap(), Gate::Proceed { .. }));
let ctx2 = ctx_with_broker(broker.clone());
let g2 = gate(
&ctx2,
shell_request("npm test"),
&[],
serde_json::json!({}),
true,
)
.await;
assert!(
matches!(g2, Gate::Proceed { .. }),
"allowlisted key should skip the prompt"
);
assert!(rx.try_recv().is_err(), "no second prompt should be sent");
}
}