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<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};
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()
}
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() }
}
}
#[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()
}
);
}
#[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");
}
#[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); 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);
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"));
}
#[test]
fn validation_invalid_parameter_types() {
#[derive(Deserialize)]
struct Params {
#[allow(unused)]
age: i32,
}
fn set_age(_: &Context, _: Params) -> Result<(), String> {
Ok(())
}
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); 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(())
}
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`"));
}
}