jsonrpce 0.1.0

JSON-RPC 2.0 for Rust
Documentation
use std::collections::HashMap;

use serde_json::Value;

use crate::Error;
use crate::Transport;
use crate::transports::BoxedHandler;
use crate::transports::Stdio;

pub struct Server<T, C>
where
    T: Transport<C>,
{
    transport: T,
    context: C,
    handlers: HashMap<String, T::Handler>,
}

// Impl for the sync transport
impl<T, C> Server<T, C>
where
    T: Transport<C, Handler = BoxedHandler<C>>,
{
    pub fn new(transport: T, context: C) -> Self {
        Server {
            transport,
            context,
            handlers: HashMap::new(),
        }
    }

    pub fn method<F, P, R, E>(mut self, name: &str, handler: F) -> Self
    where
        F: Fn(&C, P) -> Result<R, E> + Send + Sync + 'static,
        P: serde::de::DeserializeOwned + Send + Sync + 'static,
        R: serde::Serialize + Send + Sync + 'static,
        E: Into<Error> + 'static,
    {
        let wrapped = move |ctx: &C, params: Value| {
            let p: P = serde_json::from_value(params).map_err(Error::invalid_params)?;
            let r = handler(ctx, p).map_err(|e| e.into())?;
            serde_json::to_value(r).map_err(Error::internal_error)
        };

        self.handlers.insert(name.to_string(), Box::new(wrapped));
        self
    }

    pub fn run(self) -> Result<T::Output, T::Error> {
        self.transport.serve(self.context, self.handlers)
    }
}

impl<C> Server<Stdio, C>
where
    C: Send + Sync + 'static,
{
    pub fn stdio(ctx: C) -> Self {
        Self::new(Stdio::new(), ctx)
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::transports::InMemory;
    use serde::{Deserialize, Serialize};

    ///////////////////////////////////////////////////////////////////////////
    // Helpers

    /// Helper to create a standard JSON-RPC 2.0 request string
    fn make_request(method: &str, params: impl Serialize, id: i32) -> String {
        let req = serde_json::json!({
            "jsonrpc": "2.0",
            "method": method,
            "params": params,
            "id": id
        });
        req.to_string()
    }

    /// Helper to parse a JSON-RPC response string into a Result<T, Error>
    /// This mimics what a client would do.
    fn parse_result<T: serde::de::DeserializeOwned>(output: &str) -> Result<T, Error> {
        #[derive(Deserialize)]
        struct RpcResponse<T> {
            result: Option<T>,
            error: Option<Error>,
        }

        let resp: RpcResponse<T> = serde_json::from_str(output).expect("Output was not valid JSON");

        if let Some(err) = resp.error {
            Err(err)
        } else {
            Ok(resp.result.unwrap())
        }
    }

    #[derive(Debug)]
    struct Context {
        db_url: String,
    }

    impl Context {
        fn new(url: &str) -> Self {
            Self { db_url: url.into() }
        }
    }

    ///////////////////////////////////////////////////////////////////////////
    // Execution Tests

    #[test]
    fn execution_simple_request_response() {
        fn add(_: &Context, p: (i32, i32)) -> Result<i32, String> {
            Ok(p.0 + p.1)
        }

        let input = make_request("add", (10, 20), 1);

        let context = Context::new("db");
        let transport = InMemory::new(&input);
        let server = Server::new(transport, context).method("add", add);

        let output = server.run().unwrap();

        let result: i32 = parse_result(&output).unwrap();
        assert_eq!(result, 30);
    }

    #[test]
    fn execution_request_with_complex_types() {
        #[derive(Deserialize, Serialize, PartialEq, Debug)]
        struct User {
            id: i32,
            name: String,
        }

        fn create_user(_: &Context, name: String) -> Result<User, String> {
            Ok(User { id: 1, name })
        }

        let input = make_request("create", "Alice", 1);

        let context = Context::new("db");
        let transport = InMemory::new(&input);
        let server = Server::new(transport, context).method("create", create_user);

        let output = server.run().unwrap();

        let user: User = parse_result(&output).unwrap();
        assert_eq!(
            user,
            User {
                id: 1,
                name: "Alice".into()
            }
        );
    }

    ///////////////////////////////////////////////////////////////////////////
    // Context Tests

    #[test]
    fn context_handlers_can_read() {
        fn get_db(ctx: &Context, _: ()) -> Result<String, String> {
            Ok(ctx.db_url.clone())
        }

        let input = make_request("config", (), 1);
        let context = Context::new("postgres://localhost:5432");
        let transport = InMemory::new(&input);

        let server = Server::new(transport, context).method("config", get_db);
        let output = server.run().unwrap();

        let url: String = parse_result(&output).unwrap();
        assert_eq!(url, "postgres://localhost:5432");
    }

    ///////////////////////////////////////////////////////////////////////////
    // Error Tests

    #[test]
    fn errors_application_logic_error() {
        fn fail(_: &Context, _: ()) -> Result<(), String> {
            Err("Database unavailable".into())
        }

        let input = make_request("fail", (), 1);
        let context = Context::new("db");
        let transport = InMemory::new(&input);
        let server = Server::new(transport, context).method("fail", fail);

        let output = server.run().unwrap();

        let err = parse_result::<()>(&output).unwrap_err();
        assert_eq!(err.code, -32000); // Default error code for strings
        assert_eq!(err.message, "Database unavailable");
    }

    #[test]
    fn errors_custom_rpc_error() {
        fn custom_fail(_: &Context, _: ()) -> Result<(), Error> {
            Err(Error::new(1234, "Validation Failed").with_data(vec!["field_a required"]))
        }

        let input = make_request("custom", (), 1);
        let context = Context::new("db");
        let transport = InMemory::new(&input);
        let server = Server::new(transport, context).method("custom", custom_fail);

        let output = server.run().unwrap();

        let err = parse_result::<()>(&output).unwrap_err();
        assert_eq!(err.code, 1234);
        assert_eq!(err.message, "Validation Failed");
        assert_eq!(err.data, Some(serde_json::json!(["field_a required"])));
    }

    #[test]
    fn errors_method_not_found() {
        let input = make_request("unknown_method", (), 1);
        let context = Context::new("db");
        let transport = InMemory::new(&input);

        // No methods registered
        let server = Server::new(transport, context);
        let output = server.run().unwrap();

        let err = parse_result::<()>(&output).unwrap_err();
        assert_eq!(err.code, -32601);
        assert!(err.message.contains("Method not found"));
    }

    ///////////////////////////////////////////////////////////////////////////
    // Validation Tests

    #[test]
    fn validation_invalid_parameter_types() {
        #[derive(Deserialize)]
        struct Params {
            #[allow(unused)]
            age: i32,
        }

        fn set_age(_: &Context, _: Params) -> Result<(), String> {
            Ok(())
        }

        // Pass a string "too old" instead of an integer for age
        let input = make_request("set_age", serde_json::json!({ "age": "too old" }), 1);

        let context = Context::new("db");
        let transport = InMemory::new(&input);
        let server = Server::new(transport, context).method("set_age", set_age);

        let output = server.run().unwrap();

        let err = parse_result::<()>(&output).unwrap_err();
        assert_eq!(err.code, -32602); // Invalid params
        assert!(err.message.contains("Invalid params"));
    }

    #[test]
    fn validation_missing_fields() {
        #[derive(Deserialize)]
        struct Params {
            #[allow(unused)]
            required_field: String,
        }

        fn check(_: &Context, _: Params) -> Result<(), String> {
            Ok(())
        }

        // Pass empty params object
        let input = make_request("check", serde_json::json!({}), 1);

        let context = Context::new("db");
        let transport = InMemory::new(&input);
        let server = Server::new(transport, context).method("check", check);

        let output = server.run().unwrap();

        let err = parse_result::<()>(&output).unwrap_err();
        assert_eq!(err.code, -32602);
        assert!(err.message.contains("missing field `required_field`"));
    }
}