use std::{
collections::VecDeque,
sync::{
atomic::{AtomicUsize, Ordering},
Arc, Mutex,
},
time::Duration,
};
use pretty_assertions::assert_eq;
use rho_sdk::{
ApprovalContext, ApprovalDecision, ApprovalFuture, ApprovalHandler, ApprovalRequest,
CancellationToken, CapabilityRequest, CapabilitySource, PathScope, SessionId,
};
use super::{
ClassificationInput, ClassifierApprovalHandler, ClassifyFn, CONSECUTIVE_DENY_ESCALATION,
TOTAL_DENY_ESCALATION,
};
use crate::permission_classifier::ClassifierVerdict;
fn request() -> ApprovalRequest {
ApprovalRequest::new(
CapabilityRequest::write_path(
"/workspace/file.txt",
PathScope::PrimaryWorkspace,
CapabilitySource::built_in_tool("write"),
),
"approval required",
)
}
fn context_with(
history: Vec<rho_sdk::model::Message>,
cancellation: CancellationToken,
) -> ApprovalContext {
ApprovalContext::new(SessionId::new(), cancellation, history)
}
#[derive(Clone)]
struct ScriptedClassifier {
outcomes: Arc<Mutex<VecDeque<ClassifierVerdict>>>,
calls: Arc<Mutex<Vec<ApprovalRequest>>>,
histories: Arc<Mutex<Vec<Vec<rho_sdk::model::Message>>>>,
cancelled: Arc<Mutex<Vec<bool>>>,
}
impl ScriptedClassifier {
fn new(outcomes: impl IntoIterator<Item = ClassifierVerdict>) -> Self {
Self {
outcomes: Arc::new(Mutex::new(outcomes.into_iter().collect())),
calls: Arc::default(),
histories: Arc::default(),
cancelled: Arc::default(),
}
}
fn classify(&self) -> ClassifyFn {
let outcomes = Arc::clone(&self.outcomes);
let calls = Arc::clone(&self.calls);
let histories = Arc::clone(&self.histories);
let cancelled = Arc::clone(&self.cancelled);
Arc::new(move |input: ClassificationInput| {
calls.lock().unwrap().push((*input.request).clone());
histories
.lock()
.unwrap()
.push(input.request.context().history().to_vec());
cancelled
.lock()
.unwrap()
.push(input.request.context().cancellation().is_cancelled());
let outcome = outcomes
.lock()
.unwrap()
.pop_front()
.expect("scripted classifier outcome");
Box::pin(std::future::ready(outcome))
})
}
fn call_count(&self) -> usize {
self.calls.lock().unwrap().len()
}
}
#[derive(Clone)]
struct ScriptedApprovals {
decisions: Arc<Mutex<VecDeque<ApprovalDecision>>>,
requests: Arc<Mutex<Vec<ApprovalRequest>>>,
}
impl ScriptedApprovals {
fn new(decisions: impl IntoIterator<Item = ApprovalDecision>) -> Self {
Self {
decisions: Arc::new(Mutex::new(decisions.into_iter().collect())),
requests: Arc::default(),
}
}
fn request_count(&self) -> usize {
self.requests.lock().unwrap().len()
}
}
impl ApprovalHandler for ScriptedApprovals {
fn request<'a>(&'a self, request: ApprovalRequest) -> ApprovalFuture<'a> {
Box::pin(async move {
self.requests.lock().unwrap().push(request);
self.decisions
.lock()
.unwrap()
.pop_front()
.expect("scripted approval decision")
})
}
}
fn handler_with(
classifier: &ScriptedClassifier,
inner: Option<Arc<dyn ApprovalHandler>>,
) -> ClassifierApprovalHandler {
ClassifierApprovalHandler::for_tests(classifier.classify(), inner)
}
#[tokio::test]
async fn isolate_resets_deny_streak_without_sharing_counters() {
let classifier = ScriptedClassifier::new([
ClassifierVerdict::Deny {
reason: "one".into(),
},
ClassifierVerdict::Deny {
reason: "two".into(),
},
ClassifierVerdict::Deny {
reason: "three".into(),
},
ClassifierVerdict::Allow,
]);
let template = Arc::new(ClassifierApprovalHandler::for_tests(
classifier.classify(),
None,
));
let first = template.isolate();
let second = template.isolate();
for _ in 0..CONSECUTIVE_DENY_ESCALATION {
assert!(matches!(
first.request(request()).await,
ApprovalDecision::Deny { .. }
));
}
assert_eq!(second.request(request()).await, ApprovalDecision::AllowOnce);
assert_eq!(classifier.call_count(), 4);
}
#[tokio::test]
async fn approval_context_history_reaches_classifier_input() {
let classifier = ScriptedClassifier::new([ClassifierVerdict::Allow]);
let handler = ClassifierApprovalHandler::for_tests(classifier.classify(), None);
let history = vec![rho_sdk::model::Message::user_text("prior workflow context")];
assert_eq!(
handler
.request(
request().with_context(context_with(history.clone(), CancellationToken::new(),))
)
.await,
ApprovalDecision::AllowOnce
);
assert_eq!(*classifier.histories.lock().unwrap(), vec![history]);
}
#[tokio::test]
async fn approval_context_cancellation_reaches_classifier_input() {
let classifier = ScriptedClassifier::new([ClassifierVerdict::Allow]);
let handler = ClassifierApprovalHandler::for_tests(classifier.classify(), None);
let cancellation = CancellationToken::new();
cancellation.cancel();
assert_eq!(
handler
.request(request().with_context(context_with(Vec::new(), cancellation)))
.await,
ApprovalDecision::AllowOnce
);
assert_eq!(*classifier.cancelled.lock().unwrap(), vec![true]);
}
#[tokio::test]
async fn allow_returns_allow_once_and_resets_consecutive_denials() {
let classifier = ScriptedClassifier::new([
ClassifierVerdict::Deny {
reason: "too broad".into(),
},
ClassifierVerdict::Deny {
reason: "still too broad".into(),
},
ClassifierVerdict::Allow,
ClassifierVerdict::Deny {
reason: "new streak one".into(),
},
ClassifierVerdict::Deny {
reason: "new streak two".into(),
},
ClassifierVerdict::Deny {
reason: "new streak three".into(),
},
]);
let inner = Arc::new(ScriptedApprovals::new([ApprovalDecision::AllowOnce]));
let handler = handler_with(&classifier, Some(inner.clone()));
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. }
));
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. }
));
assert_eq!(
handler.request(request()).await,
ApprovalDecision::AllowOnce
);
for _ in 0..CONSECUTIVE_DENY_ESCALATION {
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. }
));
}
assert_eq!(classifier.call_count(), 6);
assert_eq!(inner.request_count(), 0);
}
#[tokio::test]
async fn after_three_denials_next_request_escalates_to_inner_handler_and_resets() {
let classifier = ScriptedClassifier::new([
ClassifierVerdict::Deny {
reason: "one".into(),
},
ClassifierVerdict::Deny {
reason: "two".into(),
},
ClassifierVerdict::Deny {
reason: "three".into(),
},
ClassifierVerdict::Deny {
reason: "after reset".into(),
},
]);
let inner = Arc::new(ScriptedApprovals::new([ApprovalDecision::AllowOnce]));
let handler = handler_with(&classifier, Some(inner.clone()));
for _ in 0..CONSECUTIVE_DENY_ESCALATION {
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. }
));
}
assert_eq!(
handler.request(request()).await,
ApprovalDecision::AllowOnce
);
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. }
));
assert_eq!(classifier.call_count(), 4);
assert_eq!(inner.request_count(), 1);
}
#[tokio::test]
async fn total_denials_escalate_to_human_and_reset_both_counters() {
let mut outcomes = Vec::new();
for index in 0..TOTAL_DENY_ESCALATION {
if index > 0 && index % 2 == 0 {
outcomes.push(ClassifierVerdict::Allow);
}
outcomes.push(ClassifierVerdict::Deny {
reason: format!("deny {index}"),
});
}
let denials_before_escalation = outcomes.len();
outcomes.push(ClassifierVerdict::Deny {
reason: "after reset".into(),
});
let classifier = ScriptedClassifier::new(outcomes);
let inner = Arc::new(ScriptedApprovals::new([ApprovalDecision::AllowOnce]));
let handler = handler_with(&classifier, Some(inner.clone()));
for _ in 0..denials_before_escalation {
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. } | ApprovalDecision::AllowOnce
));
}
assert_eq!(inner.request_count(), 0);
assert_eq!(
handler.request(request()).await,
ApprovalDecision::AllowOnce
);
assert_eq!(inner.request_count(), 1);
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. }
));
assert_eq!(inner.request_count(), 1);
assert_eq!(classifier.call_count(), denials_before_escalation + 1);
}
#[tokio::test]
async fn unavailable_denials_escalate_headless_without_further_classifier_calls() {
let classifier = ScriptedClassifier::new([
ClassifierVerdict::Deny {
reason: "classifier unavailable".into(),
},
ClassifierVerdict::Deny {
reason: "classifier unavailable".into(),
},
ClassifierVerdict::Deny {
reason: "classifier unavailable".into(),
},
]);
let handler = handler_with(&classifier, None);
for _ in 0..CONSECUTIVE_DENY_ESCALATION {
let decision = handler.request(request()).await;
let ApprovalDecision::Deny { reason } = decision else {
panic!("classifier failures must deny");
};
assert!(reason.contains("find a safer path"));
assert!(reason.contains("do not route around this block"));
}
let decision = handler.request(request()).await;
let ApprovalDecision::Deny { reason } = decision else {
panic!("headless escalation must deny");
};
assert!(reason.contains("permission classifier denied"));
assert_eq!(classifier.call_count(), 3);
}
#[tokio::test]
async fn headless_escalation_cancels_context_run_token() {
let classifier = ScriptedClassifier::new([
ClassifierVerdict::Deny {
reason: "one".into(),
},
ClassifierVerdict::Deny {
reason: "two".into(),
},
ClassifierVerdict::Deny {
reason: "three".into(),
},
]);
let handler = handler_with(&classifier, None);
let cancellation = CancellationToken::new();
for _ in 0..CONSECUTIVE_DENY_ESCALATION {
assert!(matches!(
handler
.request(request().with_context(context_with(Vec::new(), cancellation.clone(),)))
.await,
ApprovalDecision::Deny { .. }
));
assert!(!cancellation.is_cancelled());
}
let decision = handler
.request(request().with_context(context_with(Vec::new(), cancellation.clone())))
.await;
let ApprovalDecision::Deny { reason } = decision else {
panic!("headless escalation must deny");
};
assert!(reason.contains("permission classifier denied"));
assert!(cancellation.is_cancelled());
assert_eq!(classifier.call_count(), 3);
}
#[test]
fn classifier_handler_reads_live_history() {
let handler = ClassifierApprovalHandler::for_tests(
Arc::new(|_: ClassificationInput| Box::pin(async { ClassifierVerdict::Allow })),
None,
);
assert!(handler.reads_live_history());
}
#[derive(Default)]
struct HoldingHuman {
prompts: AtomicUsize,
release: tokio::sync::Notify,
}
impl ApprovalHandler for HoldingHuman {
fn request<'a>(&'a self, _request: ApprovalRequest) -> ApprovalFuture<'a> {
Box::pin(async move {
self.prompts.fetch_add(1, Ordering::SeqCst);
self.release.notified().await;
ApprovalDecision::AllowOnce
})
}
}
#[tokio::test]
async fn concurrent_escalations_prompt_the_human_once() {
let mut outcomes: Vec<_> = (0..CONSECUTIVE_DENY_ESCALATION)
.map(|index| ClassifierVerdict::Deny {
reason: format!("deny {index}"),
})
.collect();
outcomes.push(ClassifierVerdict::Allow);
let classifier = ScriptedClassifier::new(outcomes);
let human = Arc::new(HoldingHuman::default());
let handler = handler_with(&classifier, Some(human.clone()));
for _ in 0..CONSECUTIVE_DENY_ESCALATION {
assert!(matches!(
handler.request(request()).await,
ApprovalDecision::Deny { .. }
));
}
let decisions = tokio::time::timeout(Duration::from_secs(5), async {
let (first, second, ()) = tokio::join!(
handler.request(request()),
handler.request(request()),
async {
while human.prompts.load(Ordering::SeqCst) == 0 {
tokio::task::yield_now().await;
}
assert_eq!(human.prompts.load(Ordering::SeqCst), 1);
human.release.notify_one();
},
);
[first, second]
})
.await
.expect("the waiting request must not prompt the human again");
assert_eq!(
decisions,
[ApprovalDecision::AllowOnce, ApprovalDecision::AllowOnce]
);
assert_eq!(human.prompts.load(Ordering::SeqCst), 1);
assert_eq!(
classifier.call_count(),
CONSECUTIVE_DENY_ESCALATION as usize + 1
);
}
#[tokio::test]
async fn late_allow_after_concurrent_denials_spend_the_budget_escalates() {
let pending: Arc<Mutex<VecDeque<tokio::sync::oneshot::Receiver<ClassifierVerdict>>>> =
Arc::default();
let mut verdicts = Vec::new();
for _ in 0..=CONSECUTIVE_DENY_ESCALATION {
let (sender, receiver) = tokio::sync::oneshot::channel();
verdicts.push(sender);
pending.lock().unwrap().push_back(receiver);
}
let classify: ClassifyFn = {
let pending = Arc::clone(&pending);
Arc::new(move |_: ClassificationInput| {
let verdict = pending
.lock()
.unwrap()
.pop_front()
.expect("one gate per call");
Box::pin(async move { verdict.await.expect("test sends every verdict") })
})
};
let handler = Arc::new(ClassifierApprovalHandler::for_tests(classify, None));
let cancellation = CancellationToken::new();
let mut tasks = tokio::task::JoinSet::new();
for _ in 0..=CONSECUTIVE_DENY_ESCALATION {
let handler = Arc::clone(&handler);
let request = request().with_context(context_with(Vec::new(), cancellation.clone()));
tasks.spawn(async move { handler.request(request).await });
}
tokio::time::timeout(Duration::from_secs(5), async {
while !pending.lock().unwrap().is_empty() {
tokio::task::yield_now().await;
}
})
.await
.expect("every request is classifying at once");
let late_allow = verdicts.pop().unwrap();
for (index, deny) in verdicts.into_iter().enumerate() {
deny.send(ClassifierVerdict::Deny {
reason: format!("deny {index}"),
})
.unwrap();
let decision = tasks.join_next().await.unwrap().unwrap();
assert!(matches!(decision, ApprovalDecision::Deny { .. }));
}
assert!(!cancellation.is_cancelled());
late_allow.send(ClassifierVerdict::Allow).unwrap();
let decision = tasks.join_next().await.unwrap().unwrap();
assert!(matches!(decision, ApprovalDecision::Deny { .. }));
assert!(cancellation.is_cancelled());
}