use mentra::{
SessionEvent, SessionEventReceiver, SessionPermissionHandle,
error::RuntimeError,
session::{PermissionDecision, PermissionRuleScope},
};
use tokio::sync::{
broadcast::error::{RecvError, TryRecvError},
oneshot,
};
use crate::{
approval::{ApprovalAnswer, ApprovalDecision, ApprovalRequest, Approver},
event::{Event, NoticeSeverity},
run::EventSink,
};
pub(super) async fn forward_events<S: EventSink, A: Approver>(
mut receiver: SessionEventReceiver,
mut sink: S,
done: oneshot::Receiver<()>,
mut approver: A,
permissions: SessionPermissionHandle,
) -> S {
tokio::pin!(done);
let mut writing = true;
loop {
tokio::select! {
biased;
received = receiver.recv() => {
match received {
Ok(event) => {
writing = handle(&event, &mut sink, &mut approver, &permissions, writing)
.await;
}
Err(RecvError::Lagged(dropped)) => {
writing = writing && emit(&mut sink, lag_notice(dropped));
}
Err(RecvError::Closed) => return sink,
}
}
_ = &mut done => {
drain(
&mut receiver,
&mut sink,
&mut approver,
&permissions,
writing,
)
.await;
return sink;
}
}
}
}
async fn drain<S: EventSink, A: Approver>(
receiver: &mut SessionEventReceiver,
sink: &mut S,
approver: &mut A,
permissions: &SessionPermissionHandle,
mut writing: bool,
) {
loop {
match receiver.try_recv() {
Ok(event) => {
writing = handle(&event, sink, approver, permissions, writing).await;
}
Err(TryRecvError::Lagged(dropped)) => {
writing = writing && emit(sink, lag_notice(dropped));
}
Err(TryRecvError::Empty | TryRecvError::Closed) => return,
}
}
}
async fn handle<S: EventSink, A: Approver>(
event: &SessionEvent,
sink: &mut S,
approver: &mut A,
permissions: &SessionPermissionHandle,
writing: bool,
) -> bool {
let followup = resolve_if_permission(event, approver, permissions).await;
let mut writing = writing && emit_session_event(sink, event);
if let Some(notice) = followup {
writing = writing && emit(sink, notice);
}
writing
}
async fn resolve_if_permission<A: Approver>(
event: &SessionEvent,
approver: &mut A,
permissions: &SessionPermissionHandle,
) -> Option<Event> {
let SessionEvent::PermissionRequested {
request_id,
tool_call_id,
tool_name,
description,
preview,
classification,
} = event
else {
return None;
};
let answer = approver
.approve(&ApprovalRequest {
request_id: request_id.clone(),
tool_call_id: tool_call_id.clone(),
tool_name: tool_name.clone(),
description: description.clone(),
input: serde_json::from_str(preview)
.unwrap_or_else(|_| serde_json::Value::String(preview.clone())),
side_effect_level: classification.as_ref().map(|c| c.side_effect_level),
})
.await;
let refused = matches!(
answer.decision,
ApprovalDecision::Deny | ApprovalDecision::DenyForSession
);
let reason = answer.reason.clone();
let recorded = permissions.resolve_permission(request_id, permission_decision(answer));
let Err(error) = recorded else {
return None;
};
let (denied, notice) = if refused {
let denied = match reason {
Some(reason) => PermissionDecision::deny().with_reason(reason),
None => PermissionDecision::deny(),
};
(denied, unremembered_notice(tool_name, &error))
} else {
let denied = PermissionDecision::deny().with_reason(format!(
"the approval for {tool_name} could not be recorded ({error}), so the call was denied"
));
(denied, downgraded_notice(tool_name, &error))
};
permissions
.resolve_permission(request_id, denied)
.is_ok()
.then_some(notice)
}
fn permission_decision(answer: ApprovalAnswer) -> PermissionDecision {
let decision = match answer.decision {
ApprovalDecision::Allow => PermissionDecision::allow(),
ApprovalDecision::Deny => PermissionDecision::deny(),
ApprovalDecision::AllowForSession => {
PermissionDecision::allow_and_remember(PermissionRuleScope::Process)
}
ApprovalDecision::DenyForSession => {
PermissionDecision::deny_and_remember(PermissionRuleScope::Process)
}
};
match answer.reason {
Some(reason) => decision.with_reason(reason),
None => decision,
}
}
fn emit_session_event<S: EventSink>(sink: &mut S, event: &SessionEvent) -> bool {
match Event::from_session_event(event) {
Some(mapped) => emit(sink, mapped),
None => true,
}
}
fn emit<S: EventSink>(sink: &mut S, event: Event) -> bool {
sink.emit(event).is_ok()
}
fn lag_notice(dropped: u64) -> Event {
Event::Notice {
severity: NoticeSeverity::Warning,
message: format!("event stream lagged; {dropped} event(s) dropped"),
}
}
fn unremembered_notice(tool_name: &str, error: &RuntimeError) -> Event {
Event::Notice {
severity: NoticeSeverity::Warning,
message: format!(
"the refusal of {tool_name} could not be remembered ({error}); \
it applied to this call only"
),
}
}
fn downgraded_notice(tool_name: &str, error: &RuntimeError) -> Event {
Event::Notice {
severity: NoticeSeverity::Warning,
message: format!(
"the answer for {tool_name} could not be recorded ({error}); the call was denied"
),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_lag_notice_says_how_many_were_lost() {
let Event::Notice { severity, message } = lag_notice(12) else {
panic!("expected a notice");
};
assert_eq!(severity, NoticeSeverity::Warning);
assert!(message.contains("12"));
}
#[test]
fn the_store_failure_notices_say_which_half_failed() {
let error = RuntimeError::Store("disk full".to_string());
let Event::Notice { severity, message } = unremembered_notice("spawn", &error) else {
panic!("expected a notice");
};
assert_eq!(severity, NoticeSeverity::Warning);
assert!(message.contains("spawn") && message.contains("could not be remembered"));
assert!(message.contains("disk full"));
let Event::Notice { severity, message } = downgraded_notice("spawn", &error) else {
panic!("expected a notice");
};
assert_eq!(severity, NoticeSeverity::Warning);
assert!(message.contains("spawn") && message.contains("the call was denied"));
assert!(message.contains("disk full"));
}
}