use mentra::{
SessionEvent, SessionEventReceiver, SessionPermissionHandle,
session::{PermissionDecision, PermissionRuleScope},
};
use tokio::sync::{
broadcast::error::{RecvError, TryRecvError},
oneshot,
};
use crate::{
approval::{ApprovalAnswer, ApprovalDecision, ApprovalRequest, Approver},
event::{Event, NoticeSeverity},
run::{EventSink, RunUsage},
};
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, RunUsage) {
tokio::pin!(done);
let mut writing = true;
let mut usage = RunUsage::default();
loop {
tokio::select! {
biased;
received = receiver.recv() => {
match received {
Ok(event) => {
usage = usage.recording(&event);
resolve_if_permission(&event, &mut approver, &permissions).await;
writing = writing && emit_session_event(&mut sink, &event);
}
Err(RecvError::Lagged(dropped)) => {
writing = writing && emit(&mut sink, lag_notice(dropped));
}
Err(RecvError::Closed) => return (sink, usage),
}
}
_ = &mut done => {
let usage = drain(&mut receiver, &mut sink, &mut approver, &permissions, writing, usage).await;
return (sink, usage);
}
}
}
}
async fn drain<S: EventSink, A: Approver>(
receiver: &mut SessionEventReceiver,
sink: &mut S,
approver: &mut A,
permissions: &SessionPermissionHandle,
mut writing: bool,
mut usage: RunUsage,
) -> RunUsage {
loop {
match receiver.try_recv() {
Ok(event) => {
usage = usage.recording(&event);
resolve_if_permission(&event, approver, permissions).await;
writing = writing && emit_session_event(sink, &event);
}
Err(TryRecvError::Lagged(dropped)) => {
writing = writing && emit(sink, lag_notice(dropped));
}
Err(TryRecvError::Empty | TryRecvError::Closed) => return usage,
}
}
}
async fn resolve_if_permission<A: Approver>(
event: &SessionEvent,
approver: &mut A,
permissions: &SessionPermissionHandle,
) {
let SessionEvent::PermissionRequested {
request_id,
tool_call_id,
tool_name,
description,
preview,
} = event
else {
return;
};
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())),
})
.await;
let _ = permissions.resolve_permission(request_id, permission_decision(answer));
}
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::Session)
}
ApprovalDecision::DenyForSession => {
PermissionDecision::deny_and_remember(PermissionRuleScope::Session)
}
};
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"),
}
}
#[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"));
}
}