diff --git a/rust/src/jsonrpc.rs b/rust/src/jsonrpc.rs index 48e6090aed..1a590c555b 100644 --- a/rust/src/jsonrpc.rs +++ b/rust/src/jsonrpc.rs @@ -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) } } @@ -730,8 +748,7 @@ 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) => { @@ -739,6 +756,10 @@ mod tests { 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:?}"), } @@ -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::(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::(json).is_err(), + "{json}" + ); + } + } + #[test] fn request_new_sets_version() { let req = JsonRpcRequest::new(42, "test.method", None); diff --git a/rust/src/router.rs b/rust/src/router.rs index ce83d5f3ed..646f999ad3 100644 --- a/rust/src/router.rs +++ b/rust/src/router.rs @@ -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( @@ -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()) @@ -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::(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::(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" ); } diff --git a/rust/src/session.rs b/rust/src/session.rs index 82e97b6606..68cc630462 100644 --- a/rust/src/session.rs +++ b/rust/src/session.rs @@ -2268,7 +2268,7 @@ async fn handle_notification( pending_external_tools: &PendingExternalTools, ) { let dispatch_start = Instant::now(); - let event = notification.event.clone(); + let event = ¬ification.event; let event_type = event.parsed_type(); if event_type == SessionEventType::PermissionRequested { tracing::debug!( @@ -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 { diff --git a/rust/tests/session_test.rs b/rust/tests/session_test.rs index 0a8bf52d5a..0310ae380f 100644 --- a/rust/tests/session_test.rs +++ b/rust/tests/session_test.rs @@ -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() { 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]