kcode-jsonrpc-wire 0.1.0

Header-omitted JSON-RPC 2.0 wire value construction and validation
Documentation
//! Header-omitted JSON-RPC 2.0 wire values.

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(&params)?;
        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(&params)?;
    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(&params)?;
    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(_),
                ..
            })
        ));
    }
}