Skip to main content

kcode_jsonrpc_wire/
lib.rs

1//! Header-omitted JSON-RPC 2.0 wire values.
2
3use serde_json::{Map, Value};
4use std::fmt;
5
6#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
7pub struct LocalRequestId(u64);
8
9impl LocalRequestId {
10    pub const fn new(value: u64) -> Self {
11        Self(value)
12    }
13    pub const fn get(self) -> u64 {
14        self.0
15    }
16}
17
18#[derive(Clone, Debug, Eq, PartialEq, Hash)]
19pub enum PeerId {
20    Signed(i64),
21    Unsigned(u64),
22    String(String),
23}
24
25impl PeerId {
26    pub const fn signed(value: i64) -> Self {
27        Self::Signed(value)
28    }
29    pub const fn unsigned(value: u64) -> Self {
30        Self::Unsigned(value)
31    }
32    pub fn string(value: impl Into<String>) -> Self {
33        Self::String(value.into())
34    }
35    pub const fn as_i64(&self) -> Option<i64> {
36        match self {
37            Self::Signed(value) => Some(*value),
38            _ => None,
39        }
40    }
41    pub const fn as_u64(&self) -> Option<u64> {
42        match self {
43            Self::Signed(value) if *value >= 0 => Some(*value as u64),
44            Self::Unsigned(value) => Some(*value),
45            Self::String(_) | Self::Signed(_) => None,
46        }
47    }
48    pub fn to_value(&self) -> Value {
49        match self {
50            Self::Signed(value) => Value::from(*value),
51            Self::Unsigned(value) => Value::from(*value),
52            Self::String(value) => Value::String(value.clone()),
53        }
54    }
55    fn parse(value: &Value) -> Result<Self, WireError> {
56        match value {
57            Value::Number(value) => value
58                .as_i64()
59                .map(Self::Signed)
60                .or_else(|| value.as_u64().map(Self::Unsigned))
61                .ok_or(WireError::UnsupportedId),
62            Value::String(value) => Ok(Self::String(value.clone())),
63            _ => Err(WireError::UnsupportedId),
64        }
65    }
66}
67
68#[derive(Clone, Debug, PartialEq)]
69pub struct RpcError {
70    pub code: i64,
71    pub message: String,
72    pub data: Option<Value>,
73}
74
75#[derive(Clone, Debug, PartialEq)]
76pub enum Message {
77    Request {
78        id: PeerId,
79        method: String,
80        params: Value,
81    },
82    Notification {
83        method: String,
84        params: Value,
85    },
86    Response {
87        id: PeerId,
88        outcome: Result<Value, RpcError>,
89    },
90}
91
92#[derive(Clone, Debug, Eq, PartialEq)]
93pub enum WireError {
94    NotObject,
95    HeaderPresent,
96    Missing(&'static str),
97    Invalid(&'static str),
98    Ambiguous,
99    UnsupportedId,
100    UnexpectedMember(String),
101}
102
103impl fmt::Display for WireError {
104    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
105        match self {
106            Self::NotObject => f.write_str("JSON-RPC message must be an object"),
107            Self::HeaderPresent => f.write_str("header-omitted message contains jsonrpc"),
108            Self::Missing(name) => write!(f, "JSON-RPC message is missing {name}"),
109            Self::Invalid(name) => write!(f, "JSON-RPC message has invalid {name}"),
110            Self::Ambiguous => f.write_str("JSON-RPC message mixes incompatible shapes"),
111            Self::UnsupportedId => f.write_str("JSON-RPC identifier must be an integer or string"),
112            Self::UnexpectedMember(name) => write!(f, "JSON-RPC message has unexpected {name}"),
113        }
114    }
115}
116
117impl std::error::Error for WireError {}
118
119pub fn parse(value: Value) -> Result<Message, WireError> {
120    let object = value.as_object().ok_or(WireError::NotObject)?;
121    if object.contains_key("jsonrpc") {
122        return Err(WireError::HeaderPresent);
123    }
124    let method = object.contains_key("method");
125    let result = object.contains_key("result");
126    let error = object.contains_key("error");
127    if method && (result || error) || result && error {
128        return Err(WireError::Ambiguous);
129    }
130    if method {
131        let request = object.contains_key("id");
132        members(
133            object,
134            if request {
135                &["id", "method", "params"]
136            } else {
137                &["method", "params"]
138            },
139        )?;
140        let method = required_string(object, "method")?;
141        if method.is_empty() {
142            return Err(WireError::Invalid("method"));
143        }
144        let params = object.get("params").cloned().unwrap_or(Value::Null);
145        valid_params(&params)?;
146        return if request {
147            Ok(Message::Request {
148                id: PeerId::parse(object.get("id").unwrap())?,
149                method,
150                params,
151            })
152        } else {
153            Ok(Message::Notification { method, params })
154        };
155    }
156    if result || error {
157        members(
158            object,
159            if result {
160                &["id", "result"]
161            } else {
162                &["id", "error"]
163            },
164        )?;
165        let id = PeerId::parse(object.get("id").ok_or(WireError::Missing("id"))?)?;
166        return if let Some(result) = object.get("result") {
167            Ok(Message::Response {
168                id,
169                outcome: Ok(result.clone()),
170            })
171        } else {
172            Ok(Message::Response {
173                id,
174                outcome: Err(parse_error(object.get("error").unwrap())?),
175            })
176        };
177    }
178    Err(WireError::Missing("method or result/error"))
179}
180
181fn members(object: &Map<String, Value>, allowed: &[&str]) -> Result<(), WireError> {
182    object
183        .keys()
184        .find(|name| !allowed.contains(&name.as_str()))
185        .map(|name| Err(WireError::UnexpectedMember(name.clone())))
186        .unwrap_or(Ok(()))
187}
188
189fn required_string(object: &Map<String, Value>, name: &'static str) -> Result<String, WireError> {
190    object
191        .get(name)
192        .and_then(Value::as_str)
193        .map(str::to_owned)
194        .ok_or_else(|| {
195            if object.contains_key(name) {
196                WireError::Invalid(name)
197            } else {
198                WireError::Missing(name)
199            }
200        })
201}
202
203fn valid_params(params: &Value) -> Result<(), WireError> {
204    if params.is_null() || params.is_array() || params.is_object() {
205        Ok(())
206    } else {
207        Err(WireError::Invalid("params"))
208    }
209}
210
211fn parse_error(value: &Value) -> Result<RpcError, WireError> {
212    let object = value.as_object().ok_or(WireError::Invalid("error"))?;
213    members(object, &["code", "message", "data"])?;
214    let code = object
215        .get("code")
216        .and_then(Value::as_i64)
217        .ok_or(WireError::Invalid("error.code"))?;
218    let message = required_string(object, "message")?;
219    Ok(RpcError {
220        code,
221        message,
222        data: object.get("data").cloned(),
223    })
224}
225
226pub fn request(
227    id: LocalRequestId,
228    method: impl Into<String>,
229    params: Value,
230) -> Result<Value, WireError> {
231    let method = method.into();
232    if method.is_empty() {
233        return Err(WireError::Invalid("method"));
234    }
235    valid_params(&params)?;
236    Ok(object([
237        ("id".into(), Value::from(id.get())),
238        ("method".into(), Value::String(method)),
239        ("params".into(), params),
240    ]))
241}
242
243pub fn notification(method: impl Into<String>, params: Value) -> Result<Value, WireError> {
244    let method = method.into();
245    if method.is_empty() {
246        return Err(WireError::Invalid("method"));
247    }
248    valid_params(&params)?;
249    Ok(object([
250        ("method".into(), Value::String(method)),
251        ("params".into(), params),
252    ]))
253}
254
255pub fn success_response(id: PeerId, result: Value) -> Value {
256    object([("id".into(), id.to_value()), ("result".into(), result)])
257}
258
259pub fn error_response(id: PeerId, error: RpcError) -> Value {
260    let mut value = Map::new();
261    value.insert("code".into(), Value::from(error.code));
262    value.insert("message".into(), Value::String(error.message));
263    if let Some(data) = error.data {
264        value.insert("data".into(), data);
265    }
266    object([
267        ("id".into(), id.to_value()),
268        ("error".into(), Value::Object(value)),
269    ])
270}
271
272fn object<const N: usize>(members: [(String, Value); N]) -> Value {
273    Value::Object(members.into_iter().collect())
274}
275
276#[cfg(test)]
277mod tests {
278    use super::*;
279    use serde_json::json;
280
281    #[test]
282    fn ids_accept_only_integers_and_have_stable_forms() {
283        assert_eq!(
284            parse(json!({"id": 7, "method": "x"})).unwrap(),
285            Message::Request {
286                id: PeerId::Signed(7),
287                method: "x".into(),
288                params: Value::Null
289            }
290        );
291        assert_eq!(
292            parse(json!({"id": -9223372036854775808i64, "result": null})).unwrap(),
293            Message::Response {
294                id: PeerId::Signed(i64::MIN),
295                outcome: Ok(Value::Null)
296            }
297        );
298        assert_eq!(
299            parse(json!({"id": 18446744073709551615u64, "result": null})).unwrap(),
300            Message::Response {
301                id: PeerId::Unsigned(u64::MAX),
302                outcome: Ok(Value::Null)
303            }
304        );
305        assert_eq!(
306            parse(json!({"id": 1.5, "method": "x"})),
307            Err(WireError::UnsupportedId)
308        );
309        assert_eq!(PeerId::signed(4).as_u64(), Some(4));
310        assert_eq!(PeerId::signed(-1).as_u64(), None);
311    }
312
313    #[test]
314    fn rejects_extra_members_in_responses_and_errors() {
315        assert_eq!(
316            parse(json!({"id": 1, "result": null, "params": []})),
317            Err(WireError::UnexpectedMember("params".into()))
318        );
319        assert_eq!(
320            parse(json!({"id": 1, "result": null, "extra": true})),
321            Err(WireError::UnexpectedMember("extra".into()))
322        );
323        assert_eq!(
324            parse(json!({"id": 1, "error": {"code": 1, "message": "x", "extra": null}})),
325            Err(WireError::UnexpectedMember("extra".into()))
326        );
327    }
328
329    #[test]
330    fn constructors_reject_invalid_inputs_and_round_trip() {
331        assert_eq!(
332            request(LocalRequestId::new(1), "", Value::Null),
333            Err(WireError::Invalid("method"))
334        );
335        assert_eq!(
336            notification("x", json!(true)),
337            Err(WireError::Invalid("params"))
338        );
339        let call = request(LocalRequestId::new(4), "work", json!([1])).unwrap();
340        assert!(matches!(parse(call), Ok(Message::Request { .. })));
341        let notice = notification("done", Value::Null).unwrap();
342        assert!(matches!(parse(notice), Ok(Message::Notification { .. })));
343        assert!(matches!(
344            parse(success_response(PeerId::string("p"), json!(true))),
345            Ok(Message::Response { outcome: Ok(_), .. })
346        ));
347        let error = RpcError {
348            code: 1,
349            message: "no".into(),
350            data: None,
351        };
352        assert!(matches!(
353            parse(error_response(PeerId::unsigned(2), error)),
354            Ok(Message::Response {
355                outcome: Err(_),
356                ..
357            })
358        ));
359    }
360}