use std::{
collections::HashMap,
sync::{Arc, Mutex},
};
use agent_client_protocol::schema::v1::{SessionMode, SessionModeId, SessionModeState};
use basis::approval::{ApprovalAnswer, ApprovalDecision, ApprovalRequest, Approver};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ApprovalMode {
Always,
#[default]
Prompt,
Never,
}
const ALWAYS: &str = "always";
const PROMPT: &str = "prompt";
const NEVER: &str = "never";
fn mode_for(id: &str) -> Option<ApprovalMode> {
match id {
ALWAYS => Some(ApprovalMode::Always),
PROMPT => Some(ApprovalMode::Prompt),
NEVER => Some(ApprovalMode::Never),
_ => None,
}
}
fn mode_id(mode: ApprovalMode) -> SessionModeId {
SessionModeId::new(match mode {
ApprovalMode::Always => ALWAYS,
ApprovalMode::Prompt => PROMPT,
ApprovalMode::Never => NEVER,
})
}
fn describe(mode: ApprovalMode) -> SessionMode {
let (name, description) = match mode {
ApprovalMode::Always => (
"Always allow",
"Act without asking. What a confined or unattended session wants.",
),
ApprovalMode::Prompt => (
"Ask each time",
"Ask before anything that changes state outside this process.",
),
ApprovalMode::Never => (
"Read only",
"Refuse anything that changes state outside this process.",
),
};
SessionMode::new(mode_id(mode), name).description(description)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModeError {
Unknown,
NotOffered,
}
impl std::fmt::Display for ModeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Unknown => f.write_str("unknown mode"),
Self::NotOffered => {
f.write_str("this session was opened read-only and cannot change mode")
}
}
}
}
#[derive(Clone)]
pub struct SessionModes {
inner: Arc<Mutex<State>>,
}
struct State {
current: ApprovalMode,
switchable: bool,
remembered: HashMap<String, bool>,
}
impl SessionModes {
pub fn new(initial: ApprovalMode) -> Self {
Self {
inner: Arc::new(Mutex::new(State {
current: initial,
switchable: !matches!(initial, ApprovalMode::Never),
remembered: HashMap::new(),
})),
}
}
pub fn current(&self) -> ApprovalMode {
self.lock().current
}
pub fn state(&self) -> SessionModeState {
let state = self.lock();
let available = if state.switchable {
vec![
describe(ApprovalMode::Always),
describe(ApprovalMode::Prompt),
describe(ApprovalMode::Never),
]
} else {
vec![describe(state.current)]
};
SessionModeState::new(mode_id(state.current), available)
}
pub fn set(&self, id: &SessionModeId) -> Result<ApprovalMode, ModeError> {
let mode = mode_for(&id.0).ok_or(ModeError::Unknown)?;
let mut state = self.lock();
if !state.switchable && mode != state.current {
return Err(ModeError::NotOffered);
}
state.current = mode;
state.remembered.clear();
Ok(mode)
}
fn remember(&self, tool_name: &str, allow: bool) {
self.lock().remembered.insert(tool_name.to_string(), allow);
}
fn remembered(&self, tool_name: &str) -> Option<bool> {
self.lock().remembered.get(tool_name).copied()
}
fn lock(&self) -> std::sync::MutexGuard<'_, State> {
self.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
}
pub struct ModedApprover<A> {
modes: SessionModes,
inner: A,
}
impl<A> ModedApprover<A> {
pub fn new(modes: SessionModes, inner: A) -> Self {
Self { modes, inner }
}
}
#[async_trait::async_trait]
impl<A: Approver> Approver for ModedApprover<A> {
async fn approve(&mut self, request: &ApprovalRequest) -> ApprovalAnswer {
match self.modes.current() {
ApprovalMode::Always => ApprovalDecision::Allow.into(),
ApprovalMode::Never => ApprovalAnswer::new(ApprovalDecision::Deny).because(format!(
"{} changes state outside this process, and this session is set to refuse that",
request.tool_name
)),
ApprovalMode::Prompt => self.ask(request).await,
}
}
}
impl<A: Approver> ModedApprover<A> {
async fn ask(&mut self, request: &ApprovalRequest) -> ApprovalAnswer {
if let Some(allow) = self.modes.remembered(&request.tool_name) {
return if allow {
ApprovalDecision::Allow.into()
} else {
ApprovalAnswer::new(ApprovalDecision::Deny).because(format!(
"{} was refused earlier in this session, and that answer still stands",
request.tool_name
))
};
}
let answer = self.inner.approve(request).await;
match answer.decision {
ApprovalDecision::AllowForSession => {
self.modes.remember(&request.tool_name, true);
ApprovalAnswer {
decision: ApprovalDecision::Allow,
..answer
}
}
ApprovalDecision::DenyForSession => {
self.modes.remember(&request.tool_name, false);
ApprovalAnswer {
decision: ApprovalDecision::Deny,
..answer
}
}
_ => answer,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::sync::atomic::{AtomicUsize, Ordering};
struct Counting {
asked: Arc<AtomicUsize>,
answer: ApprovalDecision,
}
#[async_trait::async_trait]
impl Approver for Counting {
async fn approve(&mut self, _request: &ApprovalRequest) -> ApprovalAnswer {
self.asked.fetch_add(1, Ordering::SeqCst);
self.answer.into()
}
}
fn request(tool_name: &str) -> ApprovalRequest {
ApprovalRequest {
request_id: "r1".to_string(),
tool_call_id: "c1".to_string(),
tool_name: tool_name.to_string(),
description: "wants to write".to_string(),
input: json!({}),
}
}
fn gate(
initial: ApprovalMode,
answer: ApprovalDecision,
) -> (SessionModes, ModedApprover<Counting>, Arc<AtomicUsize>) {
let modes = SessionModes::new(initial);
let asked = Arc::new(AtomicUsize::new(0));
let approver = ModedApprover::new(
modes.clone(),
Counting {
asked: Arc::clone(&asked),
answer,
},
);
(modes, approver, asked)
}
#[test]
fn every_offered_mode_maps_back_to_one_lan_can_read() {
for mode in SessionModes::new(ApprovalMode::Prompt)
.state()
.available_modes
{
assert!(
mode_for(&mode.id.0).is_some(),
"offered {} but cannot read it back",
mode.id.0
);
}
}
#[test]
fn the_state_reports_the_current_mode_and_all_three() {
let state = SessionModes::new(ApprovalMode::Prompt).state();
assert_eq!(&*state.current_mode_id.0, PROMPT);
assert_eq!(state.available_modes.len(), 3);
}
#[test]
fn a_read_only_session_offers_nothing_else() {
let modes = SessionModes::new(ApprovalMode::Never);
let state = modes.state();
assert_eq!(state.available_modes.len(), 1);
assert_eq!(&*state.current_mode_id.0, NEVER);
assert_eq!(
modes.set(&SessionModeId::new(ALWAYS)),
Err(ModeError::NotOffered),
"a client cannot lift a prohibition it was never given"
);
}
#[test]
fn switching_reports_the_new_mode() {
let modes = SessionModes::new(ApprovalMode::Prompt);
assert_eq!(
modes.set(&SessionModeId::new(ALWAYS)),
Ok(ApprovalMode::Always)
);
assert_eq!(modes.current(), ApprovalMode::Always);
}
#[test]
fn an_unknown_mode_is_refused() {
let modes = SessionModes::new(ApprovalMode::Prompt);
assert_eq!(
modes.set(&SessionModeId::new("architect")),
Err(ModeError::Unknown)
);
assert_eq!(
modes.current(),
ApprovalMode::Prompt,
"a refused switch must leave the session where it was"
);
}
#[tokio::test]
async fn allow_and_refuse_answer_without_asking() {
for (mode, expected) in [
(ApprovalMode::Always, ApprovalDecision::Allow),
(ApprovalMode::Never, ApprovalDecision::Deny),
] {
let (_modes, mut approver, asked) = gate(mode, ApprovalDecision::Allow);
assert_eq!(approver.approve(&request("shell")).await.decision, expected);
assert_eq!(
asked.load(Ordering::SeqCst),
0,
"{mode:?} has nothing to ask about"
);
}
}
#[tokio::test]
async fn a_read_only_session_says_so_when_it_refuses() {
let (_modes, mut approver, _asked) = gate(ApprovalMode::Never, ApprovalDecision::Allow);
assert_eq!(
approver.approve(&request("shell")).await.reason.as_deref(),
Some(
"shell changes state outside this process, \
and this session is set to refuse that"
)
);
}
#[tokio::test]
async fn asking_puts_the_request_to_the_client() {
let (_modes, mut approver, asked) = gate(ApprovalMode::Prompt, ApprovalDecision::Allow);
assert_eq!(
approver.approve(&request("shell")).await.decision,
ApprovalDecision::Allow
);
assert_eq!(asked.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn an_answer_for_the_session_is_not_asked_twice() {
let (_modes, mut approver, asked) =
gate(ApprovalMode::Prompt, ApprovalDecision::AllowForSession);
assert_eq!(
approver.approve(&request("shell")).await.decision,
ApprovalDecision::Allow
);
assert_eq!(
approver.approve(&request("shell")).await.decision,
ApprovalDecision::Allow
);
assert_eq!(asked.load(Ordering::SeqCst), 1);
assert_eq!(
approver.approve(&request("files")).await.decision,
ApprovalDecision::Allow
);
assert_eq!(asked.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn changing_mode_forgets_what_was_allowed_for_the_session() {
let (modes, mut approver, _asked) =
gate(ApprovalMode::Prompt, ApprovalDecision::AllowForSession);
approver.approve(&request("shell")).await;
modes.set(&SessionModeId::new(NEVER)).expect("switches");
assert_eq!(
approver.approve(&request("shell")).await.decision,
ApprovalDecision::Deny,
"a stale allow must not survive the mode that replaced it"
);
modes.set(&SessionModeId::new(PROMPT)).expect("switches");
assert_eq!(
approver.approve(&request("shell")).await.decision,
ApprovalDecision::Allow,
"the client is asked again, and answered again"
);
}
#[tokio::test]
async fn a_refusal_for_the_session_is_also_remembered() {
let (_modes, mut approver, asked) =
gate(ApprovalMode::Prompt, ApprovalDecision::DenyForSession);
assert_eq!(
approver.approve(&request("shell")).await.decision,
ApprovalDecision::Deny
);
let repeated = approver.approve(&request("shell")).await;
assert_eq!(repeated.decision, ApprovalDecision::Deny);
assert_eq!(
repeated.reason.as_deref(),
Some("shell was refused earlier in this session, and that answer still stands"),
"a remembered refusal still owes the model a reason"
);
assert_eq!(asked.load(Ordering::SeqCst), 1);
}
}