11// Copyright (c) Microsoft Corporation.
22// Licensed under the MIT License.
33
4- use crate :: send_error;
4+ use crate :: { send_error, RequestId } ;
55use serde_json:: { self , Value } ;
66use 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 ) > ;
1313type 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
1616pub 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