Skip to content

Commit 110daaa

Browse files
karthiknadigCopilot
andcommitted
fix: preserve JSONRPC request IDs (Fixes #549)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent 22e37af commit 110daaa

6 files changed

Lines changed: 389 additions & 105 deletions

File tree

‎crates/pet-jsonrpc/src/lib.rs‎

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,20 @@ use std::io::{self, Write};
66

77
pub mod server;
88

9+
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
10+
#[serde(untagged)]
11+
pub enum RequestId {
12+
String(String),
13+
Number(serde_json::Number),
14+
Null,
15+
}
16+
17+
impl From<u32> for RequestId {
18+
fn from(id: u32) -> Self {
19+
Self::Number(id.into())
20+
}
21+
}
22+
923
#[derive(Serialize, Deserialize)]
1024
#[serde(rename_all = "camelCase")]
1125
#[derive(Debug)]
@@ -29,7 +43,7 @@ pub fn send_message<T: serde::Serialize>(method: &'static str, params: Option<T>
2943
);
3044
let _ = io::stdout().flush();
3145
}
32-
pub fn send_reply<T: serde::Serialize>(id: u32, payload: Option<T>) {
46+
pub fn send_reply<T: serde::Serialize>(id: &RequestId, payload: Option<T>) {
3347
let payload = serde_json::json!({
3448
"jsonrpc": "2.0",
3549
"result": payload,
@@ -44,7 +58,7 @@ pub fn send_reply<T: serde::Serialize>(id: u32, payload: Option<T>) {
4458
let _ = io::stdout().flush();
4559
}
4660

47-
pub fn send_error(id: Option<u32>, code: i32, message: String) {
61+
pub fn send_error(id: Option<&RequestId>, code: i32, message: String) {
4862
let payload = serde_json::json!({
4963
"jsonrpc": "2.0",
5064
"error": { "code": code, "message": message },

‎crates/pet-jsonrpc/src/server.rs‎

Lines changed: 97 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,17 @@
11
// Copyright (c) Microsoft Corporation.
22
// Licensed under the MIT License.
33

4-
use crate::send_error;
4+
use crate::{send_error, RequestId};
55
use serde_json::{self, Value};
66
use std::{
77
collections::HashMap,
88
io::{self, Read},
99
sync::Arc,
1010
};
1111

12-
type RequestHandler<C> = Arc<dyn Fn(Arc<C>, u32, Value)>;
12+
type RequestHandler<C> = Arc<dyn Fn(Arc<C>, RequestId, Value)>;
1313
type NotificationHandler<C> = Arc<dyn Fn(Arc<C>, Value)>;
14-
type ErrorHandler = Arc<dyn Fn(Option<u32>, i32, String)>;
14+
type ErrorHandler = Arc<dyn Fn(Option<RequestId>, i32, String)>;
1515

1616
pub struct HandlersKeyedByMethodName<C> {
1717
context: Arc<C>,
@@ -26,14 +26,14 @@ impl<C> HandlersKeyedByMethodName<C> {
2626
context,
2727
requests: HashMap::new(),
2828
notifications: HashMap::new(),
29-
send_error: Arc::new(send_error),
29+
send_error: Arc::new(|id, code, message| send_error(id.as_ref(), code, message)),
3030
}
3131
}
3232

3333
#[cfg(test)]
3434
fn new_with_error_handler(
3535
context: Arc<C>,
36-
send_error: impl Fn(Option<u32>, i32, String) + 'static,
36+
send_error: impl Fn(Option<RequestId>, i32, String) + 'static,
3737
) -> Self {
3838
HandlersKeyedByMethodName {
3939
context,
@@ -45,7 +45,7 @@ impl<C> HandlersKeyedByMethodName<C> {
4545

4646
pub fn add_request_handler<F>(&mut self, method: &'static str, handler: F)
4747
where
48-
F: Fn(Arc<C>, u32, Value) + Send + Sync + 'static,
48+
F: Fn(Arc<C>, RequestId, Value) + Send + Sync + 'static,
4949
{
5050
self.requests.insert(
5151
method,
@@ -68,32 +68,39 @@ impl<C> HandlersKeyedByMethodName<C> {
6868
}
6969

7070
fn handle_request(&self, message: Value) {
71+
let id = match message.get("id") {
72+
None => None,
73+
Some(Value::String(id)) => Some(RequestId::String(id.clone())),
74+
Some(Value::Number(id)) => Some(RequestId::Number(id.clone())),
75+
Some(Value::Null) => Some(RequestId::Null),
76+
Some(_) => {
77+
(self.send_error)(None, -32600, "Invalid JSONRPC request ID".to_string());
78+
return;
79+
}
80+
};
7181
match message["method"].as_str() {
7282
Some(method) => {
73-
if let Some(id) = message["id"].as_u64() {
83+
if let Some(id) = id {
7484
if let Some(handler) = self.requests.get(method) {
75-
handler(self.context.clone(), id as u32, message["params"].clone());
85+
handler(self.context.clone(), id, message["params"].clone());
7686
} else {
7787
eprint!("Failed to find handler for method: {method}");
7888
(self.send_error)(
79-
Some(id as u32),
89+
Some(id),
8090
-1,
8191
format!("Failed to find handler for request {method}"),
8292
);
8393
}
94+
} else if let Some(handler) = self.notifications.get(method) {
95+
handler(self.context.clone(), message["params"].clone());
8496
} else {
85-
// No id, so this is a notification
86-
if let Some(handler) = self.notifications.get(method) {
87-
handler(self.context.clone(), message["params"].clone());
88-
} else {
89-
eprint!("Failed to find handler for method: {method}");
90-
}
97+
eprint!("Failed to find handler for method: {method}");
9198
}
9299
}
93100
None => {
94101
eprint!("Failed to get method from message: {message}");
95102
(self.send_error)(
96-
message["id"].as_u64().map(|id| id as u32),
103+
id,
97104
-3,
98105
format!("Failed to extract method from JSONRPC payload {message:?}"),
99106
);
@@ -176,9 +183,9 @@ mod tests {
176183

177184
#[derive(Default)]
178185
struct TestContext {
179-
request: Mutex<Option<(u32, Value)>>,
186+
request: Mutex<Option<(RequestId, Value)>>,
180187
notification: Mutex<Option<Value>>,
181-
errors: Mutex<Vec<(Option<u32>, i32, String)>>,
188+
errors: Mutex<Vec<(Option<RequestId>, i32, String)>>,
182189
}
183190

184191
fn create_handlers_with_recorded_errors(
@@ -194,6 +201,74 @@ mod tests {
194201
})
195202
}
196203

204+
#[test]
205+
fn request_ids_preserve_values_for_dispatch_and_errors() {
206+
for value in [
207+
json!("request-1"),
208+
json!(""),
209+
json!("\u{03c0}-request"),
210+
json!(-7),
211+
json!(0),
212+
json!(u64::from(u32::MAX) + 1),
213+
json!(u64::MAX),
214+
json!(i64::MIN),
215+
json!(1.5),
216+
Value::Null,
217+
] {
218+
let id = serde_json::from_value::<RequestId>(value.clone()).unwrap();
219+
assert_eq!(serde_json::to_value(&id).unwrap(), value);
220+
let context = Arc::new(TestContext::default());
221+
let mut handlers = create_handlers_with_recorded_errors(context.clone());
222+
handlers.add_request_handler("method", |context, id, params| {
223+
*context.request.lock().unwrap() = Some((id, params));
224+
});
225+
handlers.add_notification_handler("method", |context, params| {
226+
*context.notification.lock().unwrap() = Some(params);
227+
});
228+
handlers.handle_request(json!({"id": value, "method": "method", "params": 42}));
229+
assert_eq!(
230+
*context.request.lock().unwrap(),
231+
Some((id.clone(), json!(42)))
232+
);
233+
assert!(context.notification.lock().unwrap().is_none());
234+
handlers.handle_request(json!({"id": value, "method": "unknown"}));
235+
handlers.handle_request(json!({"id": value}));
236+
let errors = context.errors.lock().unwrap();
237+
assert_eq!(errors.len(), 2);
238+
assert_eq!(errors[0].0, Some(id.clone()));
239+
assert_eq!(errors[0].1, -1);
240+
assert_eq!(errors[1].0, Some(id));
241+
assert_eq!(errors[1].1, -3);
242+
}
243+
}
244+
245+
#[test]
246+
fn missing_id_is_a_notification_but_invalid_ids_are_rejected() {
247+
let context = Arc::new(TestContext::default());
248+
let mut handlers = create_handlers_with_recorded_errors(context.clone());
249+
handlers.add_request_handler("method", |context, id, params| {
250+
*context.request.lock().unwrap() = Some((id, params));
251+
});
252+
handlers.add_notification_handler("method", |context, params| {
253+
*context.notification.lock().unwrap() = Some(params);
254+
});
255+
handlers.handle_request(json!({"method": "method", "params": 7}));
256+
assert_eq!(context.notification.lock().unwrap().take(), Some(json!(7)));
257+
assert!(context.request.lock().unwrap().is_none());
258+
assert!(context.errors.lock().unwrap().is_empty());
259+
for value in [json!(true), json!(false), json!([]), json!({"id": 1})] {
260+
assert!(serde_json::from_value::<RequestId>(value.clone()).is_err());
261+
handlers.handle_request(json!({"id": value, "method": "method"}));
262+
assert!(context.notification.lock().unwrap().is_none());
263+
assert!(context.request.lock().unwrap().is_none());
264+
assert_eq!(
265+
context.errors.lock().unwrap().pop(),
266+
Some((None, -32600, "Invalid JSONRPC request ID".into()))
267+
);
268+
}
269+
assert!(context.errors.lock().unwrap().is_empty());
270+
}
271+
197272
#[test]
198273
fn get_content_length_parses_valid_header() {
199274
assert_eq!(get_content_length("Content-Length: 42\r\n").unwrap(), 42);
@@ -238,7 +313,7 @@ mod tests {
238313

239314
assert_eq!(
240315
*context.request.lock().unwrap(),
241-
Some((7, json!({ "value": 42 })))
316+
Some((7.into(), json!({ "value": 42 })))
242317
);
243318
assert_eq!(*context.notification.lock().unwrap(), Some(json!(["item"])));
244319
}
@@ -259,7 +334,7 @@ mod tests {
259334

260335
assert_eq!(
261336
*context.request.lock().unwrap(),
262-
Some((9, json!({ "ok": true })))
337+
Some((9.into(), json!({ "ok": true })))
263338
);
264339
}
265340

@@ -309,7 +384,7 @@ mod tests {
309384
assert_eq!(
310385
context.errors.lock().unwrap().as_slice(),
311386
&[(
312-
Some(1),
387+
Some(1.into()),
313388
-1,
314389
"Failed to find handler for request unknown/request".to_string()
315390
)]
@@ -332,7 +407,7 @@ mod tests {
332407
assert_eq!(
333408
context.errors.lock().unwrap().as_slice(),
334409
&[(
335-
Some(1),
410+
Some(1.into()),
336411
-3,
337412
format!("Failed to extract method from JSONRPC payload {message:?}")
338413
)]

0 commit comments

Comments
 (0)