use serde_json::{Map, Value};
use std::fmt;
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub struct LocalRequestId(u64);
impl LocalRequestId {
pub const fn new(value: u64) -> Self {
Self(value)
}
pub const fn get(self) -> u64 {
self.0
}
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub enum PeerId {
Signed(i64),
Unsigned(u64),
String(String),
}
impl PeerId {
pub const fn signed(value: i64) -> Self {
Self::Signed(value)
}
pub const fn unsigned(value: u64) -> Self {
Self::Unsigned(value)
}
pub fn string(value: impl Into<String>) -> Self {
Self::String(value.into())
}
pub const fn as_i64(&self) -> Option<i64> {
match self {
Self::Signed(value) => Some(*value),
_ => None,
}
}
pub const fn as_u64(&self) -> Option<u64> {
match self {
Self::Signed(value) if *value >= 0 => Some(*value as u64),
Self::Unsigned(value) => Some(*value),
Self::String(_) | Self::Signed(_) => None,
}
}
pub fn to_value(&self) -> Value {
match self {
Self::Signed(value) => Value::from(*value),
Self::Unsigned(value) => Value::from(*value),
Self::String(value) => Value::String(value.clone()),
}
}
fn parse(value: &Value) -> Result<Self, WireError> {
match value {
Value::Number(value) => value
.as_i64()
.map(Self::Signed)
.or_else(|| value.as_u64().map(Self::Unsigned))
.ok_or(WireError::UnsupportedId),
Value::String(value) => Ok(Self::String(value.clone())),
_ => Err(WireError::UnsupportedId),
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct RpcError {
pub code: i64,
pub message: String,
pub data: Option<Value>,
}
#[derive(Clone, Debug, PartialEq)]
pub enum Message {
Request {
id: PeerId,
method: String,
params: Value,
},
Notification {
method: String,
params: Value,
},
Response {
id: PeerId,
outcome: Result<Value, RpcError>,
},
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum WireError {
NotObject,
HeaderPresent,
Missing(&'static str),
Invalid(&'static str),
Ambiguous,
UnsupportedId,
UnexpectedMember(String),
}
impl fmt::Display for WireError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::NotObject => f.write_str("JSON-RPC message must be an object"),
Self::HeaderPresent => f.write_str("header-omitted message contains jsonrpc"),
Self::Missing(name) => write!(f, "JSON-RPC message is missing {name}"),
Self::Invalid(name) => write!(f, "JSON-RPC message has invalid {name}"),
Self::Ambiguous => f.write_str("JSON-RPC message mixes incompatible shapes"),
Self::UnsupportedId => f.write_str("JSON-RPC identifier must be an integer or string"),
Self::UnexpectedMember(name) => write!(f, "JSON-RPC message has unexpected {name}"),
}
}
}
impl std::error::Error for WireError {}
pub fn parse(value: Value) -> Result<Message, WireError> {
let object = value.as_object().ok_or(WireError::NotObject)?;
if object.contains_key("jsonrpc") {
return Err(WireError::HeaderPresent);
}
let method = object.contains_key("method");
let result = object.contains_key("result");
let error = object.contains_key("error");
if method && (result || error) || result && error {
return Err(WireError::Ambiguous);
}
if method {
let request = object.contains_key("id");
members(
object,
if request {
&["id", "method", "params"]
} else {
&["method", "params"]
},
)?;
let method = required_string(object, "method")?;
if method.is_empty() {
return Err(WireError::Invalid("method"));
}
let params = object.get("params").cloned().unwrap_or(Value::Null);
valid_params(¶ms)?;
return if request {
Ok(Message::Request {
id: PeerId::parse(object.get("id").unwrap())?,
method,
params,
})
} else {
Ok(Message::Notification { method, params })
};
}
if result || error {
members(
object,
if result {
&["id", "result"]
} else {
&["id", "error"]
},
)?;
let id = PeerId::parse(object.get("id").ok_or(WireError::Missing("id"))?)?;
return if let Some(result) = object.get("result") {
Ok(Message::Response {
id,
outcome: Ok(result.clone()),
})
} else {
Ok(Message::Response {
id,
outcome: Err(parse_error(object.get("error").unwrap())?),
})
};
}
Err(WireError::Missing("method or result/error"))
}
fn members(object: &Map<String, Value>, allowed: &[&str]) -> Result<(), WireError> {
object
.keys()
.find(|name| !allowed.contains(&name.as_str()))
.map(|name| Err(WireError::UnexpectedMember(name.clone())))
.unwrap_or(Ok(()))
}
fn required_string(object: &Map<String, Value>, name: &'static str) -> Result<String, WireError> {
object
.get(name)
.and_then(Value::as_str)
.map(str::to_owned)
.ok_or_else(|| {
if object.contains_key(name) {
WireError::Invalid(name)
} else {
WireError::Missing(name)
}
})
}
fn valid_params(params: &Value) -> Result<(), WireError> {
if params.is_null() || params.is_array() || params.is_object() {
Ok(())
} else {
Err(WireError::Invalid("params"))
}
}
fn parse_error(value: &Value) -> Result<RpcError, WireError> {
let object = value.as_object().ok_or(WireError::Invalid("error"))?;
members(object, &["code", "message", "data"])?;
let code = object
.get("code")
.and_then(Value::as_i64)
.ok_or(WireError::Invalid("error.code"))?;
let message = required_string(object, "message")?;
Ok(RpcError {
code,
message,
data: object.get("data").cloned(),
})
}
pub fn request(
id: LocalRequestId,
method: impl Into<String>,
params: Value,
) -> Result<Value, WireError> {
let method = method.into();
if method.is_empty() {
return Err(WireError::Invalid("method"));
}
valid_params(¶ms)?;
Ok(object([
("id".into(), Value::from(id.get())),
("method".into(), Value::String(method)),
("params".into(), params),
]))
}
pub fn notification(method: impl Into<String>, params: Value) -> Result<Value, WireError> {
let method = method.into();
if method.is_empty() {
return Err(WireError::Invalid("method"));
}
valid_params(¶ms)?;
Ok(object([
("method".into(), Value::String(method)),
("params".into(), params),
]))
}
pub fn success_response(id: PeerId, result: Value) -> Value {
object([("id".into(), id.to_value()), ("result".into(), result)])
}
pub fn error_response(id: PeerId, error: RpcError) -> Value {
let mut value = Map::new();
value.insert("code".into(), Value::from(error.code));
value.insert("message".into(), Value::String(error.message));
if let Some(data) = error.data {
value.insert("data".into(), data);
}
object([
("id".into(), id.to_value()),
("error".into(), Value::Object(value)),
])
}
fn object<const N: usize>(members: [(String, Value); N]) -> Value {
Value::Object(members.into_iter().collect())
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn ids_accept_only_integers_and_have_stable_forms() {
assert_eq!(
parse(json!({"id": 7, "method": "x"})).unwrap(),
Message::Request {
id: PeerId::Signed(7),
method: "x".into(),
params: Value::Null
}
);
assert_eq!(
parse(json!({"id": -9223372036854775808i64, "result": null})).unwrap(),
Message::Response {
id: PeerId::Signed(i64::MIN),
outcome: Ok(Value::Null)
}
);
assert_eq!(
parse(json!({"id": 18446744073709551615u64, "result": null})).unwrap(),
Message::Response {
id: PeerId::Unsigned(u64::MAX),
outcome: Ok(Value::Null)
}
);
assert_eq!(
parse(json!({"id": 1.5, "method": "x"})),
Err(WireError::UnsupportedId)
);
assert_eq!(PeerId::signed(4).as_u64(), Some(4));
assert_eq!(PeerId::signed(-1).as_u64(), None);
}
#[test]
fn rejects_extra_members_in_responses_and_errors() {
assert_eq!(
parse(json!({"id": 1, "result": null, "params": []})),
Err(WireError::UnexpectedMember("params".into()))
);
assert_eq!(
parse(json!({"id": 1, "result": null, "extra": true})),
Err(WireError::UnexpectedMember("extra".into()))
);
assert_eq!(
parse(json!({"id": 1, "error": {"code": 1, "message": "x", "extra": null}})),
Err(WireError::UnexpectedMember("extra".into()))
);
}
#[test]
fn constructors_reject_invalid_inputs_and_round_trip() {
assert_eq!(
request(LocalRequestId::new(1), "", Value::Null),
Err(WireError::Invalid("method"))
);
assert_eq!(
notification("x", json!(true)),
Err(WireError::Invalid("params"))
);
let call = request(LocalRequestId::new(4), "work", json!([1])).unwrap();
assert!(matches!(parse(call), Ok(Message::Request { .. })));
let notice = notification("done", Value::Null).unwrap();
assert!(matches!(parse(notice), Ok(Message::Notification { .. })));
assert!(matches!(
parse(success_response(PeerId::string("p"), json!(true))),
Ok(Message::Response { outcome: Ok(_), .. })
));
let error = RpcError {
code: 1,
message: "no".into(),
data: None,
};
assert!(matches!(
parse(error_response(PeerId::unsigned(2), error)),
Ok(Message::Response {
outcome: Err(_),
..
})
));
}
}