1use std::{
42 collections::HashMap,
43 sync::{Arc, Mutex},
44};
45
46use agent_client_protocol::schema::v1::{SessionMode, SessionModeId, SessionModeState};
47
48use basis::approval::{ApprovalAnswer, ApprovalDecision, ApprovalRequest, Approver};
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
56pub enum ApprovalMode {
57 Always,
59 #[default]
63 Prompt,
64 Never,
67}
68
69const ALWAYS: &str = "always";
73const PROMPT: &str = "prompt";
74const NEVER: &str = "never";
75
76fn mode_for(id: &str) -> Option<ApprovalMode> {
78 match id {
79 ALWAYS => Some(ApprovalMode::Always),
80 PROMPT => Some(ApprovalMode::Prompt),
81 NEVER => Some(ApprovalMode::Never),
82 _ => None,
83 }
84}
85
86fn mode_id(mode: ApprovalMode) -> SessionModeId {
87 SessionModeId::new(match mode {
88 ApprovalMode::Always => ALWAYS,
89 ApprovalMode::Prompt => PROMPT,
90 ApprovalMode::Never => NEVER,
91 })
92}
93
94fn describe(mode: ApprovalMode) -> SessionMode {
96 let (name, description) = match mode {
97 ApprovalMode::Always => (
98 "Always allow",
99 "Act without asking. What a confined or unattended session wants.",
100 ),
101 ApprovalMode::Prompt => (
102 "Ask each time",
103 "Ask before anything that changes state outside this process.",
104 ),
105 ApprovalMode::Never => (
106 "Read only",
107 "Refuse anything that changes state outside this process.",
108 ),
109 };
110
111 SessionMode::new(mode_id(mode), name).description(description)
112}
113
114#[derive(Debug, Clone, Copy, PartialEq, Eq)]
116pub enum ModeError {
117 Unknown,
119 NotOffered,
121}
122
123impl std::fmt::Display for ModeError {
124 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
125 match self {
126 Self::Unknown => f.write_str("unknown mode"),
127 Self::NotOffered => {
128 f.write_str("this session was opened read-only and cannot change mode")
129 }
130 }
131 }
132}
133
134#[derive(Clone)]
141pub struct SessionModes {
142 inner: Arc<Mutex<State>>,
143}
144
145struct State {
146 current: ApprovalMode,
147 switchable: bool,
149 remembered: HashMap<String, bool>,
151}
152
153impl SessionModes {
154 pub fn new(initial: ApprovalMode) -> Self {
163 Self {
164 inner: Arc::new(Mutex::new(State {
165 current: initial,
166 switchable: !matches!(initial, ApprovalMode::Never),
167 remembered: HashMap::new(),
168 })),
169 }
170 }
171
172 pub fn current(&self) -> ApprovalMode {
173 self.lock().current
174 }
175
176 pub fn state(&self) -> SessionModeState {
179 let state = self.lock();
180 let available = if state.switchable {
181 vec![
182 describe(ApprovalMode::Always),
183 describe(ApprovalMode::Prompt),
184 describe(ApprovalMode::Never),
185 ]
186 } else {
187 vec![describe(state.current)]
188 };
189
190 SessionModeState::new(mode_id(state.current), available)
191 }
192
193 pub fn set(&self, id: &SessionModeId) -> Result<ApprovalMode, ModeError> {
200 let mode = mode_for(&id.0).ok_or(ModeError::Unknown)?;
201
202 let mut state = self.lock();
203 if !state.switchable && mode != state.current {
204 return Err(ModeError::NotOffered);
205 }
206
207 state.current = mode;
208 state.remembered.clear();
209 Ok(mode)
210 }
211
212 fn remember(&self, tool_name: &str, allow: bool) {
213 self.lock().remembered.insert(tool_name.to_string(), allow);
214 }
215
216 fn remembered(&self, tool_name: &str) -> Option<bool> {
217 self.lock().remembered.get(tool_name).copied()
218 }
219
220 fn lock(&self) -> std::sync::MutexGuard<'_, State> {
221 self.inner
225 .lock()
226 .unwrap_or_else(|poisoned| poisoned.into_inner())
227 }
228}
229
230pub struct ModedApprover<A> {
233 modes: SessionModes,
234 inner: A,
235}
236
237impl<A> ModedApprover<A> {
238 pub fn new(modes: SessionModes, inner: A) -> Self {
239 Self { modes, inner }
240 }
241}
242
243#[async_trait::async_trait]
244impl<A: Approver> Approver for ModedApprover<A> {
245 async fn approve(&mut self, request: &ApprovalRequest) -> ApprovalAnswer {
246 match self.modes.current() {
250 ApprovalMode::Always => ApprovalDecision::Allow.into(),
251 ApprovalMode::Never => ApprovalAnswer::new(ApprovalDecision::Deny).because(format!(
252 "{} changes state outside this process, and this session is set to refuse that",
253 request.tool_name
254 )),
255 ApprovalMode::Prompt => self.ask(request).await,
256 }
257 }
258}
259
260impl<A: Approver> ModedApprover<A> {
261 async fn ask(&mut self, request: &ApprovalRequest) -> ApprovalAnswer {
262 if let Some(allow) = self.modes.remembered(&request.tool_name) {
263 return if allow {
264 ApprovalDecision::Allow.into()
265 } else {
266 ApprovalAnswer::new(ApprovalDecision::Deny).because(format!(
267 "{} was refused earlier in this session, and that answer still stands",
268 request.tool_name
269 ))
270 };
271 }
272
273 let answer = self.inner.approve(request).await;
278 match answer.decision {
279 ApprovalDecision::AllowForSession => {
280 self.modes.remember(&request.tool_name, true);
281 ApprovalAnswer {
282 decision: ApprovalDecision::Allow,
283 ..answer
284 }
285 }
286 ApprovalDecision::DenyForSession => {
287 self.modes.remember(&request.tool_name, false);
288 ApprovalAnswer {
289 decision: ApprovalDecision::Deny,
290 ..answer
291 }
292 }
293 _ => answer,
294 }
295 }
296}
297
298#[cfg(test)]
299mod tests {
300 use super::*;
301 use basis::ToolSideEffectLevel;
302 use serde_json::json;
303 use std::sync::atomic::{AtomicUsize, Ordering};
304
305 struct Counting {
308 asked: Arc<AtomicUsize>,
309 answer: ApprovalDecision,
310 }
311
312 #[async_trait::async_trait]
313 impl Approver for Counting {
314 async fn approve(&mut self, _request: &ApprovalRequest) -> ApprovalAnswer {
315 self.asked.fetch_add(1, Ordering::SeqCst);
316 self.answer.into()
317 }
318 }
319
320 fn request(tool_name: &str) -> ApprovalRequest {
321 ApprovalRequest {
322 request_id: "r1".to_string(),
323 tool_call_id: "c1".to_string(),
324 tool_name: tool_name.to_string(),
325 description: "wants to write".to_string(),
326 input: json!({}),
327 side_effect_level: Some(ToolSideEffectLevel::LocalState),
328 }
329 }
330
331 fn gate(
332 initial: ApprovalMode,
333 answer: ApprovalDecision,
334 ) -> (SessionModes, ModedApprover<Counting>, Arc<AtomicUsize>) {
335 let modes = SessionModes::new(initial);
336 let asked = Arc::new(AtomicUsize::new(0));
337 let approver = ModedApprover::new(
338 modes.clone(),
339 Counting {
340 asked: Arc::clone(&asked),
341 answer,
342 },
343 );
344 (modes, approver, asked)
345 }
346
347 #[test]
348 fn every_offered_mode_maps_back_to_one_lan_can_read() {
349 for mode in SessionModes::new(ApprovalMode::Prompt)
352 .state()
353 .available_modes
354 {
355 assert!(
356 mode_for(&mode.id.0).is_some(),
357 "offered {} but cannot read it back",
358 mode.id.0
359 );
360 }
361 }
362
363 #[test]
364 fn the_state_reports_the_current_mode_and_all_three() {
365 let state = SessionModes::new(ApprovalMode::Prompt).state();
366
367 assert_eq!(&*state.current_mode_id.0, PROMPT);
368 assert_eq!(state.available_modes.len(), 3);
369 }
370
371 #[test]
372 fn a_read_only_session_offers_nothing_else() {
373 let modes = SessionModes::new(ApprovalMode::Never);
374 let state = modes.state();
375
376 assert_eq!(state.available_modes.len(), 1);
377 assert_eq!(&*state.current_mode_id.0, NEVER);
378 assert_eq!(
379 modes.set(&SessionModeId::new(ALWAYS)),
380 Err(ModeError::NotOffered),
381 "a client cannot lift a prohibition it was never given"
382 );
383 }
384
385 #[test]
386 fn switching_reports_the_new_mode() {
387 let modes = SessionModes::new(ApprovalMode::Prompt);
388
389 assert_eq!(
390 modes.set(&SessionModeId::new(ALWAYS)),
391 Ok(ApprovalMode::Always)
392 );
393 assert_eq!(modes.current(), ApprovalMode::Always);
394 }
395
396 #[test]
397 fn an_unknown_mode_is_refused() {
398 let modes = SessionModes::new(ApprovalMode::Prompt);
399
400 assert_eq!(
401 modes.set(&SessionModeId::new("architect")),
402 Err(ModeError::Unknown)
403 );
404 assert_eq!(
405 modes.current(),
406 ApprovalMode::Prompt,
407 "a refused switch must leave the session where it was"
408 );
409 }
410
411 #[tokio::test]
412 async fn allow_and_refuse_answer_without_asking() {
413 for (mode, expected) in [
414 (ApprovalMode::Always, ApprovalDecision::Allow),
415 (ApprovalMode::Never, ApprovalDecision::Deny),
416 ] {
417 let (_modes, mut approver, asked) = gate(mode, ApprovalDecision::Allow);
418
419 assert_eq!(approver.approve(&request("shell")).await.decision, expected);
420 assert_eq!(
421 asked.load(Ordering::SeqCst),
422 0,
423 "{mode:?} has nothing to ask about"
424 );
425 }
426 }
427
428 #[tokio::test]
429 async fn a_read_only_session_says_so_when_it_refuses() {
430 let (_modes, mut approver, _asked) = gate(ApprovalMode::Never, ApprovalDecision::Allow);
433
434 assert_eq!(
435 approver.approve(&request("shell")).await.reason.as_deref(),
436 Some(
437 "shell changes state outside this process, \
438 and this session is set to refuse that"
439 )
440 );
441 }
442
443 #[tokio::test]
444 async fn asking_puts_the_request_to_the_client() {
445 let (_modes, mut approver, asked) = gate(ApprovalMode::Prompt, ApprovalDecision::Allow);
446
447 assert_eq!(
448 approver.approve(&request("shell")).await.decision,
449 ApprovalDecision::Allow
450 );
451 assert_eq!(asked.load(Ordering::SeqCst), 1);
452 }
453
454 #[tokio::test]
455 async fn an_answer_for_the_session_is_not_asked_twice() {
456 let (_modes, mut approver, asked) =
457 gate(ApprovalMode::Prompt, ApprovalDecision::AllowForSession);
458
459 assert_eq!(
462 approver.approve(&request("shell")).await.decision,
463 ApprovalDecision::Allow
464 );
465 assert_eq!(
466 approver.approve(&request("shell")).await.decision,
467 ApprovalDecision::Allow
468 );
469 assert_eq!(asked.load(Ordering::SeqCst), 1);
470
471 assert_eq!(
473 approver.approve(&request("files")).await.decision,
474 ApprovalDecision::Allow
475 );
476 assert_eq!(asked.load(Ordering::SeqCst), 2);
477 }
478
479 #[tokio::test]
480 async fn changing_mode_forgets_what_was_allowed_for_the_session() {
481 let (modes, mut approver, _asked) =
482 gate(ApprovalMode::Prompt, ApprovalDecision::AllowForSession);
483
484 approver.approve(&request("shell")).await;
485 modes.set(&SessionModeId::new(NEVER)).expect("switches");
486
487 assert_eq!(
488 approver.approve(&request("shell")).await.decision,
489 ApprovalDecision::Deny,
490 "a stale allow must not survive the mode that replaced it"
491 );
492
493 modes.set(&SessionModeId::new(PROMPT)).expect("switches");
495 assert_eq!(
496 approver.approve(&request("shell")).await.decision,
497 ApprovalDecision::Allow,
498 "the client is asked again, and answered again"
499 );
500 }
501
502 #[tokio::test]
503 async fn a_refusal_for_the_session_is_also_remembered() {
504 let (_modes, mut approver, asked) =
505 gate(ApprovalMode::Prompt, ApprovalDecision::DenyForSession);
506
507 assert_eq!(
508 approver.approve(&request("shell")).await.decision,
509 ApprovalDecision::Deny
510 );
511
512 let repeated = approver.approve(&request("shell")).await;
513 assert_eq!(repeated.decision, ApprovalDecision::Deny);
514 assert_eq!(
515 repeated.reason.as_deref(),
516 Some("shell was refused earlier in this session, and that answer still stands"),
517 "a remembered refusal still owes the model a reason"
518 );
519 assert_eq!(asked.load(Ordering::SeqCst), 1);
520 }
521}