use std::collections::HashSet;
use serde_json::Value;
use tokio::sync::mpsc;
use crate::channels::{AgentUpdate, ApprovalDecision, ApprovalRequest, CancelFlag};
use crate::risk::{self, Capability, Decision};
static SESSION: std::sync::OnceLock<std::sync::Arc<Approver>> = std::sync::OnceLock::new();
pub fn install(approver: std::sync::Arc<Approver>) {
let _ = SESSION.set(approver);
}
pub fn session() -> Option<std::sync::Arc<Approver>> {
SESSION.get().cloned()
}
pub struct Approver {
updates: mpsc::UnboundedSender<AgentUpdate>,
decisions: tokio::sync::Mutex<mpsc::UnboundedReceiver<ApprovalDecision>>,
allowed: std::sync::Mutex<HashSet<String>>,
cancel: CancelFlag,
unattended: Option<Unattended>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Unattended {
Refuse,
Proceed,
}
impl Approver {
pub fn new(
updates: mpsc::UnboundedSender<AgentUpdate>,
decisions: mpsc::UnboundedReceiver<ApprovalDecision>,
cancel: CancelFlag,
) -> Self {
Self {
updates,
decisions: tokio::sync::Mutex::new(decisions),
allowed: std::sync::Mutex::new(HashSet::new()),
cancel,
unattended: None,
}
}
pub fn unattended(allow_changes: bool) -> Self {
let (updates, _updates_rx) = mpsc::unbounded_channel();
let (_decisions_tx, decisions) = mpsc::unbounded_channel();
Self {
updates,
decisions: tokio::sync::Mutex::new(decisions),
allowed: std::sync::Mutex::new(HashSet::new()),
cancel: CancelFlag::default(),
unattended: Some(if allow_changes {
Unattended::Proceed
} else {
Unattended::Refuse
}),
}
}
pub async fn approve(
&self,
tool: &str,
capability: Capability,
input: &Value,
) -> Result<(), String> {
let assessment = risk::assess(tool, capability, input);
match assessment.decision {
Decision::Allow => return Ok(()),
Decision::Deny(reason) => return Err(reason),
Decision::Confirm => {}
}
match self.unattended {
Some(Unattended::Proceed) => return Ok(()),
Some(Unattended::Refuse) => {
return Err(format!(
"{} needs approval and this run has no interface to ask. Nothing ran. Start \
the run with --allow-changes if it is meant to change things.",
assessment.detail
))
}
None => {}
}
if self
.allowed
.lock()
.map(|set| set.contains(&assessment.scope))
.unwrap_or(false)
{
return Ok(());
}
if self.cancel.is_raised() {
return Err("The user interrupted the turn, so this tool did not run.".to_string());
}
let mut decisions = self.decisions.lock().await;
if self
.updates
.send(AgentUpdate::Approval(ApprovalRequest {
tool: tool.to_string(),
detail: assessment.detail,
scope: assessment.scope.clone(),
}))
.is_err()
{
return Err(
"Could not ask the user for approval, so this tool did not run.".to_string(),
);
}
let decision = tokio::select! {
answer = decisions.recv() => answer,
_ = self.cancel.wait() => None,
};
match decision {
Some(ApprovalDecision::Once) => Ok(()),
Some(ApprovalDecision::Always) => {
if let Ok(mut set) = self.allowed.lock() {
set.insert(assessment.scope);
}
Ok(())
}
Some(ApprovalDecision::Deny) => Err(format!(
"The user declined to let {} run. Do not retry it; ask what they want instead.",
tool
)),
None => Err("The tool was not approved, so it did not run.".to_string()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn approver() -> (
Approver,
mpsc::UnboundedReceiver<AgentUpdate>,
mpsc::UnboundedSender<ApprovalDecision>,
CancelFlag,
) {
let (updates, updates_rx) = mpsc::unbounded_channel();
let (decisions_tx, decisions) = mpsc::unbounded_channel();
let cancel = CancelFlag::default();
(
Approver::new(updates, decisions, cancel.clone()),
updates_rx,
decisions_tx,
cancel,
)
}
#[tokio::test]
async fn a_read_is_never_put_to_the_user() {
let (approver, mut updates, _tx, _cancel) = approver();
assert!(approver
.approve("read_file", Capability::ReadOnly, &json!({}))
.await
.is_ok());
assert!(
updates.try_recv().is_err(),
"a read must not raise a prompt"
);
}
#[tokio::test]
async fn approving_once_does_not_approve_the_next_call() {
let (approver, mut updates, tx, _cancel) = approver();
let call = json!({"path": "a.rs"});
tx.send(ApprovalDecision::Once).unwrap();
assert!(approver
.approve("write_file", Capability::Write, &call)
.await
.is_ok());
assert!(updates.try_recv().is_ok(), "the first call must ask");
tx.send(ApprovalDecision::Once).unwrap();
assert!(approver
.approve("write_file", Capability::Write, &call)
.await
.is_ok());
assert!(
updates.try_recv().is_ok(),
"'once' means once; the second call must ask again"
);
}
#[tokio::test]
async fn always_stops_asking_for_that_scope_only() {
let (approver, mut updates, tx, _cancel) = approver();
let lib = json!({"path": "src/lib.rs"});
tx.send(ApprovalDecision::Always).unwrap();
assert!(approver
.approve("write_file", Capability::Write, &lib)
.await
.is_ok());
let _ = updates.try_recv();
assert!(approver
.approve("write_file", Capability::Write, &lib)
.await
.is_ok());
assert!(
updates.try_recv().is_err(),
"'always' must stop the prompt for the scope it was granted for"
);
tx.send(ApprovalDecision::Deny).unwrap();
assert!(approver
.approve(
"write_file",
Capability::Write,
&json!({"path": "src/main.rs"})
)
.await
.is_err());
tx.send(ApprovalDecision::Deny).unwrap();
assert!(approver
.approve(
"caatinga_deploy",
Capability::Signing,
&json!({"network": "testnet"})
)
.await
.is_err());
}
#[tokio::test]
async fn a_denial_tells_the_model_not_to_retry() {
let (approver, _updates, tx, _cancel) = approver();
tx.send(ApprovalDecision::Deny).unwrap();
let err = approver
.approve("write_file", Capability::Write, &json!({"path": "a.rs"}))
.await
.unwrap_err();
assert!(err.contains("Do not retry"), "got {}", err);
}
#[tokio::test]
async fn a_refusal_is_not_put_to_the_user_as_a_question() {
let (approver, mut updates, _tx, _cancel) = approver();
let err = approver
.approve(
"stellar_invoke",
Capability::Signing,
&json!({"source": "S".repeat(56)}),
)
.await
.unwrap_err();
assert!(err.contains("identity alias"), "got {}", err);
assert!(
updates.try_recv().is_err(),
"a denial must not raise a prompt"
);
}
#[tokio::test]
async fn an_unattended_run_refuses_rather_than_waits() {
let approver = Approver::unattended(false);
let err = approver
.approve("write_file", Capability::Write, &json!({"path": "a.rs"}))
.await
.unwrap_err();
assert!(err.contains("a.rs"), "got {}", err);
assert!(err.contains("--allow-changes"), "got {}", err);
}
#[tokio::test]
async fn an_unattended_run_told_to_proceed_proceeds() {
let approver = Approver::unattended(true);
assert!(approver
.approve("write_file", Capability::Write, &json!({"path": "a.rs"}))
.await
.is_ok());
}
#[tokio::test]
async fn an_unattended_yes_does_not_override_a_refusal() {
let approver = Approver::unattended(true);
let err = approver
.approve(
"stellar_invoke",
Capability::Signing,
&json!({"source": "S".repeat(56)}),
)
.await
.unwrap_err();
assert!(err.contains("identity alias"), "got {}", err);
}
#[tokio::test]
async fn cancelling_releases_a_turn_parked_on_a_question() {
let (approver, _updates, _tx, cancel) = approver();
let raise = tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(80)).await;
cancel.raise();
});
let result = approver
.approve("write_file", Capability::Write, &json!({}))
.await;
raise.await.unwrap();
assert!(result.is_err(), "a cancelled question must not approve");
}
#[tokio::test]
async fn a_closed_channel_refuses_rather_than_proceeds() {
let (approver, _updates, tx, _cancel) = approver();
drop(tx);
assert!(approver
.approve("write_file", Capability::Write, &json!({}))
.await
.is_err());
}
}