1use 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(¶ms)?;
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(¶ms)?;
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(¶ms)?;
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}