use std::collections::BTreeMap;
use std::marker::PhantomData;
use std::sync::Arc;
use async_trait::async_trait;
use serde::Serialize;
use serde::de::DeserializeOwned;
use super::wire::{
CODE_INVALID_REQUEST, CODE_PARSE_ERROR, JSONRPC_VERSION, RpcError, RpcRequest, RpcResponse,
};
#[async_trait]
pub trait RpcMethod: Send + Sync + 'static {
async fn call(&self, params: serde_json::Value) -> Result<serde_json::Value, RpcError>;
}
struct Typed<Req, Resp, F> {
call: F,
_types: PhantomData<fn(Req) -> Resp>,
}
#[async_trait]
impl<Req, Resp, F, Fut> RpcMethod for Typed<Req, Resp, F>
where
Req: DeserializeOwned + Send + 'static,
Resp: Serialize + Send + 'static,
F: Fn(Req) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = Result<Resp, RpcError>> + Send + 'static,
{
async fn call(&self, params: serde_json::Value) -> Result<serde_json::Value, RpcError> {
let request: Req = serde_json::from_value(params)
.map_err(|e| RpcError::invalid_params(format!("params do not decode: {e}")))?;
let response = (self.call)(request).await?;
serde_json::to_value(&response)
.map_err(|e| RpcError::internal(format!("serialize response: {e}")))
}
}
pub fn typed_method<Req, Resp, F, Fut>(call: F) -> impl RpcMethod
where
Req: DeserializeOwned + Send + 'static,
Resp: Serialize + Send + 'static,
F: Fn(Req) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = Result<Resp, RpcError>> + Send + 'static,
{
Typed::<Req, Resp, F> {
call,
_types: PhantomData,
}
}
#[derive(Default)]
pub struct RpcRouter {
methods: BTreeMap<String, Arc<dyn RpcMethod>>,
}
impl std::fmt::Debug for RpcRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RpcRouter")
.field("methods", &self.method_names().collect::<Vec<_>>())
.finish()
}
}
impl RpcRouter {
pub fn new() -> Self {
Self::default()
}
pub fn method(mut self, name: impl Into<String>, handler: impl RpcMethod) -> Self {
self.methods.insert(name.into(), Arc::new(handler));
self
}
pub fn typed<Req, Resp, F, Fut>(self, name: impl Into<String>, call: F) -> Self
where
Req: DeserializeOwned + Send + 'static,
Resp: Serialize + Send + 'static,
F: Fn(Req) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = Result<Resp, RpcError>> + Send + 'static,
{
self.method(name, typed_method::<Req, Resp, F, Fut>(call))
}
pub fn method_names(&self) -> impl Iterator<Item = &str> {
self.methods.keys().map(String::as_str)
}
pub async fn dispatch(&self, frame: &[u8]) -> RpcResponse {
let request: RpcRequest = match serde_json::from_slice(frame) {
Ok(r) => r,
Err(e) => {
return RpcResponse::failure(
serde_json::Value::Null,
RpcError::new(CODE_PARSE_ERROR, format!("unparseable request frame: {e}")),
);
}
};
if request.jsonrpc != JSONRPC_VERSION {
return RpcResponse::failure(
request.id,
RpcError::new(
CODE_INVALID_REQUEST,
format!(
"unsupported jsonrpc version {:?}; this listener speaks {JSONRPC_VERSION}",
request.jsonrpc
),
),
);
}
let Some(handler) = self.methods.get(&request.method) else {
let known: Vec<&str> = self.method_names().collect();
return RpcResponse::failure(
request.id,
RpcError::method_not_found(&request.method, &known),
);
};
let handler = Arc::clone(handler);
match handler.call(request.params).await {
Ok(result) => RpcResponse::success(request.id, result),
Err(error) => RpcResponse::failure(request.id, error),
}
}
}