Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 80 additions & 19 deletions rust/src/router.rs
Original file line number Diff line number Diff line change
Expand Up @@ -204,18 +204,15 @@ impl SessionRouter {
let Some(mut params) = notification.params else {
continue;
};
let Some(session_id) = params.get("sessionId").and_then(|v| v.as_str())
let Some(session_id) = params
.get("sessionId")
.and_then(|v| v.as_str())
.map(str::to_owned)
else {
continue;
};

let sender = {
let guard = sessions.lock();
guard
.get_key_value(session_id)
.map(|(id, s)| (id.clone(), s.notifications.clone()))
};
if let Some((session_id, sender)) = sender {
if sessions.lock().contains_key(session_id.as_str()) {
// Leave null in the existing slot so serde still rejects
// missing data, without rebuilding the owned payload.
let data = params
Expand All @@ -227,7 +224,7 @@ impl SessionRouter {
if let Some(data) = data {
event_notification.event.data = data;
}
let _ = sender.send(event_notification);
route_notification(&sessions, &session_id, event_notification);
}
Err(e) => {
warn!(
Expand Down Expand Up @@ -289,17 +286,12 @@ impl SessionRouter {
.and_then(|p| p.get("sessionId"))
.and_then(|v| v.as_str());

if let Some(sid) = session_id {
let sender = {
let guard = sessions.lock();
guard.get(sid).map(|s| s.requests.clone())
};
if let Some(sender) = sender {
let _ = sender.send(request);
} else {
if let Some(sid) = session_id.map(str::to_owned) {
let method = request.method.clone();
if !route_request(&sessions, &sid, request) {
warn!(
session_id = sid,
method = %request.method,
session_id = %sid,
method = %method,
"request for unregistered session"
);
}
Expand All @@ -315,6 +307,32 @@ impl SessionRouter {
}
}

// Keep lookup and enqueue in one critical section so a replacement registration
// cannot receive a message through a sender captured from its predecessor.
fn route_notification(
sessions: &Mutex<HashMap<SessionId, SessionSenders>>,
session_id: &str,
notification: SessionEventNotification,
) {
let guard = sessions.lock();
if let Some(senders) = guard.get(session_id) {
let _ = senders.notifications.send(notification);
}
}

fn route_request(
sessions: &Mutex<HashMap<SessionId, SessionSenders>>,
session_id: &str,
request: JsonRpcRequest,
) -> bool {
let guard = sessions.lock();
let Some(senders) = guard.get(session_id) else {
return false;
};
let _ = senders.requests.send(request);
true
}

#[cfg(test)]
mod tests {
use super::*;
Expand Down Expand Up @@ -347,4 +365,47 @@ mod tests {
// Removing twice is a no-op rather than evicting a future tenant.
assert!(!router.unregister_owned(&session_id(), live.token));
}

#[test]
fn routed_requests_and_notifications_reach_the_current_registration() {
let router = SessionRouter::new();
let mut stale = router.register(&session_id());
let mut current = router.register(&session_id());

assert!(route_request(
&router.sessions,
session_id().as_str(),
JsonRpcRequest::new(1, "tool.call", None),
));
assert_eq!(
current.channels.requests.try_recv().unwrap().method,
"tool.call"
);

route_notification(
&router.sessions,
session_id().as_str(),
SessionEventNotification {
session_id: session_id(),
event: crate::types::SessionEvent {
id: "event-1".to_string(),
timestamp: "2026-01-01T00:00:00Z".to_string(),
parent_id: None,
ephemeral: None,
agent_id: None,
debug_cli_received_at_ms: None,
debug_ws_forwarded_at_ms: None,
event_type: "session.idle".to_string(),
data: serde_json::json!({}),
},
},
);
assert_eq!(
current.channels.notifications.try_recv().unwrap().event.id,
"event-1"
);

assert!(stale.channels.requests.try_recv().is_err());
assert!(stale.channels.notifications.try_recv().is_err());
}
}
Loading