Skip to content
Open
Show file tree
Hide file tree
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
88 changes: 81 additions & 7 deletions rust/src/jsonrpc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -124,25 +124,43 @@ impl<'de> Deserialize<'de> for JsonRpcMessage {
where
D: serde::Deserializer<'de>,
{
let value = Value::deserialize(deserializer)?;
let mut value = Value::deserialize(deserializer)?;
let obj = value
.as_object()
.as_object_mut()
.ok_or_else(|| serde::de::Error::custom("expected a JSON object"))?;

let has_id = obj.contains_key("id");
let has_method = obj.contains_key("method");

// Preserve the owned payload instead of rebuilding its JSON containers
// while serde validates the envelope. Optional null payloads remain None.
let payload_key = if has_id && !has_method {
"result"
} else {
"params"
};
let payload = obj.remove(payload_key).filter(|value| !value.is_null());

if has_id && has_method {
JsonRpcRequest::deserialize(value)
.map(JsonRpcMessage::Request)
.map(|mut request| {
request.params = payload;
JsonRpcMessage::Request(request)
})
.map_err(serde::de::Error::custom)
} else if has_id {
JsonRpcResponse::deserialize(value)
.map(JsonRpcMessage::Response)
.map(|mut response| {
response.result = payload;
JsonRpcMessage::Response(response)
})
.map_err(serde::de::Error::custom)
} else {
JsonRpcNotification::deserialize(value)
.map(JsonRpcMessage::Notification)
.map(|mut notification| {
notification.params = payload;
JsonRpcMessage::Notification(notification)
})
.map_err(serde::de::Error::custom)
}
}
Expand Down Expand Up @@ -730,15 +748,18 @@ mod tests {

#[test]
fn deserialize_error_response() {
let json =
r#"{"jsonrpc":"2.0","id":7,"error":{"code":-32600,"message":"Invalid Request"}}"#;
let json = r#"{"jsonrpc":"2.0","id":7,"error":{"code":-32600,"message":"Invalid Request","data":{"nested":[1,{"reason":"invalid"}]}}}"#;
let msg: JsonRpcMessage = serde_json::from_str(json).unwrap();
match msg {
JsonRpcMessage::Response(r) => {
assert!(r.is_error());
let err = r.error.unwrap();
assert_eq!(err.code, -32600);
assert_eq!(err.message, "Invalid Request");
assert_eq!(
err.data,
Some(serde_json::json!({"nested": [1, {"reason": "invalid"}]}))
);
}
other => panic!("expected Response, got {other:?}"),
}
Expand All @@ -750,6 +771,59 @@ mod tests {
assert!(result.is_err());
}

#[test]
fn deserialize_preserves_optional_payloads() {
for payload in [
None,
Some(Value::Null),
Some(serde_json::json!(false)),
Some(serde_json::json!(42)),
Some(serde_json::json!("text")),
Some(serde_json::json!([{"nested": [1, null, true]}])),
Some(serde_json::json!({"rows": [{"content": "result"}]})),
] {
for mut envelope in [
serde_json::json!({"jsonrpc": "2.0", "method": "notify"}),
serde_json::json!({"jsonrpc": "2.0", "id": 1, "method": "request"}),
serde_json::json!({"jsonrpc": "2.0", "id": 1}),
] {
let (payload_key, ignored_key) = if envelope.get("method").is_some() {
("params", "result")
} else {
("result", "params")
};
envelope[ignored_key] = serde_json::json!({"ignored": "opposite payload"});
if let Some(payload) = &payload {
envelope[payload_key] = payload.clone();
}
let actual = match serde_json::from_value::<JsonRpcMessage>(envelope).unwrap() {
JsonRpcMessage::Request(request) => request.params,
JsonRpcMessage::Response(response) => response.result,
JsonRpcMessage::Notification(notification) => notification.params,
};
assert_eq!(actual, payload.clone().filter(|value| !value.is_null()));
}
}
}

#[test]
fn deserialize_rejects_invalid_metadata() {
for json in [
r#"{"jsonrpc":null,"method":"notify","params":{"nested":[1]}}"#,
r#"{"jsonrpc":"2.0","method":42,"params":{"nested":[1]}}"#,
r#"{"jsonrpc":"2.0","id":null,"result":{}}"#,
r#"{"jsonrpc":"2.0","id":"1","result":{}}"#,
r#"{"jsonrpc":"2.0","id":-1,"result":{}}"#,
r#"{"jsonrpc":"2.0","id":1,"method":null,"params":{}}"#,
r#"{"jsonrpc":"2.0","id":1,"result":{},"error":{"code":"bad","message":"error"}}"#,
] {
assert!(
serde_json::from_str::<JsonRpcMessage>(json).is_err(),
"{json}"
);
}
}

#[test]
fn request_new_sets_version() {
let req = JsonRpcRequest::new(42, "test.method", None);
Expand Down
28 changes: 19 additions & 9 deletions rust/src/router.rs
Original file line number Diff line number Diff line change
Expand Up @@ -174,12 +174,12 @@ impl SessionRouter {
// callback (if any) registered at client construction.
if notification.method == "gitHubTelemetry.event" {
if let Some(ref callback) = github_telemetry {
let Some(ref params) = notification.params else {
let Some(params) = notification.params else {
continue;
};
match serde_json::from_value::<
crate::github_telemetry::GitHubTelemetryNotification,
>(params.clone())
>(params)
{
Ok(telemetry) => {
if std::panic::catch_unwind(std::panic::AssertUnwindSafe(
Expand All @@ -206,7 +206,7 @@ impl SessionRouter {
if notification.method != "session.event" {
continue;
}
let Some(ref params) = notification.params else {
let Some(mut params) = notification.params else {
continue;
};
let Some(session_id) = params.get("sessionId").and_then(|v| v.as_str())
Expand All @@ -216,18 +216,28 @@ impl SessionRouter {

let sender = {
let guard = sessions.lock();
guard.get(session_id).map(|s| s.notifications.clone())
guard
.get_key_value(session_id)
.map(|(id, s)| (id.clone(), s.notifications.clone()))
};
if let Some(sender) = sender {
match serde_json::from_value::<SessionEventNotification>(params.clone())
{
Ok(event_notification) => {
if let Some((session_id, sender)) = sender {
// Leave null in the existing slot so serde still rejects
// missing data, without rebuilding the owned payload.
let data = params
.get_mut("event")
.and_then(|event| event.get_mut("data"))
.map(serde_json::Value::take);
match serde_json::from_value::<SessionEventNotification>(params) {
Ok(mut event_notification) => {
if let Some(data) = data {
event_notification.event.data = data;
}
let _ = sender.send(event_notification);
}
Err(e) => {
warn!(
error = %e,
session_id = session_id,
session_id = %session_id,
"failed to deserialize session event notification"
);
}
Expand Down
4 changes: 2 additions & 2 deletions rust/src/session.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2268,7 +2268,7 @@ async fn handle_notification(
pending_external_tools: &PendingExternalTools,
) {
let dispatch_start = Instant::now();
let event = notification.event.clone();
let event = &notification.event;
Comment thread
mohamedmansour marked this conversation as resolved.
let event_type = event.parsed_type();
if event_type == SessionEventType::PermissionRequested {
tracing::debug!(
Expand Down Expand Up @@ -2298,7 +2298,7 @@ async fn handle_notification(
}
waiter.last_assistant_message = Some(event.clone());
}
SessionEventType::SessionIdle if is_autopilot_continuation_idle(&event) => {}
SessionEventType::SessionIdle if is_autopilot_continuation_idle(event) => {}
SessionEventType::SessionIdle | SessionEventType::SessionError => {
if let Some(waiter) = guard.take() {
if event_type == SessionEventType::SessionIdle {
Expand Down
53 changes: 48 additions & 5 deletions rust/tests/session_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3846,15 +3846,58 @@ async fn permission_result_forwards_context_beside_result() {
}

#[tokio::test]
async fn session_event_notification_reaches_handler() {
async fn session_event_notification_reaches_subscribers() {
Comment thread
mohamedmansour marked this conversation as resolved.
let (session, mut server) = create_session_pair().await;
let mut sub = session.subscribe();
let mut first = session.subscribe();
let mut second = session.subscribe();
let data = serde_json::json!({"deltaContent": "hello"});

server
.send_event("session.idle", serde_json::json!({}))
.send_notification(
"session.event",
serde_json::json!({"sessionId": server.session_id, "event": null}),
)
.await;
server
.send_event("assistant.message_delta", data.clone())
.await;

let event = timeout(TIMEOUT, sub.recv()).await.unwrap().unwrap();
assert_eq!(event.event_type, "session.idle");
let mut first_event = timeout(TIMEOUT, first.recv()).await.unwrap().unwrap();
assert_eq!(first_event.data, data);
first_event.data["deltaContent"] = serde_json::json!("changed");

let second_event = timeout(TIMEOUT, second.recv()).await.unwrap().unwrap();
assert_eq!(second_event.id, first_event.id);
assert_eq!(second_event.event_type, "assistant.message_delta");
assert_eq!(second_event.data, data);
}

#[tokio::test]
async fn session_event_notification_preserves_unknown_event_payloads() {
let (session, mut server) = create_session_pair().await;
let mut events = session.subscribe();
server
.send_notification(
"session.event",
serde_json::json!({
"sessionId": server.session_id,
"event": {"id": "invalid", "timestamp": "now", "type": "future.event"}
}),
)
.await;

for data in [
Value::Null,
serde_json::json!("text"),
serde_json::json!(false),
serde_json::json!([1, {"nested": [null, true]}]),
serde_json::json!({"result": {"rows": [{"content": "preserved"}]}}),
] {
server.send_event("future.event", data.clone()).await;
let event = timeout(TIMEOUT, events.recv()).await.unwrap().unwrap();
assert_eq!(event.event_type, "future.event");
assert_eq!(event.data, data);
}
}

#[tokio::test]
Expand Down