pub mod ctx;
pub mod marker;
pub mod response;
pub mod ts;
use futures::{Stream, StreamExt};
use jsonrpsee::{
DisconnectError, RpcModule, SubscriptionCloseResponse, SubscriptionMessage,
types::{Params, ResponsePayload},
};
use serde::Deserialize;
use serde_json::json;
use ts_rs::TS;
use std::pin::pin;
use self::{ctx::FromRequestExtensions, response::ResponseValue, ts::TsTypeTuple};
pub trait QubitHandler<Ctx, MSig>: 'static + Send + Sync + Clone {
type Ctx: 'static + Send + Sync + FromRequestExtensions<Ctx>;
type Params: TsTypeTuple;
type Return;
fn call(&self, ctx: Self::Ctx, params: Params) -> Self::Return;
}
macro_rules! impl_handlers {
(impl [$($ctx:ident, $($params:ident,)*)?]) => {
impl<Ctx, F, R, $($ctx, $($params),*)?> QubitHandler<
Ctx,
(
($($ctx, $($params,)*)?),
R
)
>
for F
where
F: 'static + Send + Sync + Clone + Fn($($ctx, $($params),*)?) -> R,
impl_handlers!(ctx_ty [$($ctx)?]): 'static + Send + Sync + FromRequestExtensions<Ctx>,
$($($params: 'static + TS + Send + for<'a> Deserialize<'a>),*)?
{
type Ctx = impl_handlers!(ctx_ty [$($ctx)?]);
type Params = ($($($params,)*)?);
type Return = R;
fn call(
&self,
#[allow(unused)] ctx: Self::Ctx,
#[allow(unused)] params: Params
) -> Self::Return {
#[allow(non_snake_case)]
let ($($($params,)*)?) = match impl_handlers!(parse_impl params -> [$($($params,)*)?]) {
Ok(params) => params,
Err(e) => {
dbg!(e);
panic!("fukc");
}
};
self($(ctx, $($params,)*)?)
}
}
};
(ctx_ty [$ctx:ty]) => {
$ctx
};
(ctx_ty []) => {
Ctx
};
(parse_impl $params:ident -> []) => {
$params.parse::<[(); 0]>()
.map(|_| ())
};
(parse_impl $params:ident -> [$($param_tys:ident,)*]) => {
$params.parse::<Self::Params>()
};
(count []) => { 0 };
(count [$param:ident, $($params:ident,)*]) => {
1 + impl_handlers!(count [$($params,)*])
};
(recurse []) => {};
(recurse [$param:ident, $($params:ident,)*]) => {
impl_handlers!($($params),*);
};
($($params:ident),* $(,)?) => {
impl_handlers!(impl [$($params,)*]);
impl_handlers!(recurse [$($params,)*]);
};
}
impl_handlers!(
P0, P1, P2, P3, P4, P5, P6, P7, P8, P9, P10, P11, P12, P13, P14, P15
);
pub trait RegisterableHandler<
Ctx,
MSig,
MValue: marker::ResponseMarker,
MReturn: marker::HandlerReturnMarker,
>: QubitHandler<Ctx, MSig>
{
type Response: ResponseValue<MValue>;
fn register(self, module: &mut RpcModule<Ctx>, method_name: String);
}
impl<Ctx, T, MSig, MValue> RegisterableHandler<Ctx, MSig, MValue, marker::MResponse<MValue>> for T
where
Ctx: 'static + Clone + Send + Sync,
MValue: marker::ResponseMarker,
T: QubitHandler<Ctx, MSig>,
T::Return: ResponseValue<MValue>,
{
type Response = T::Return;
fn register(self, module: &mut RpcModule<Ctx>, method_name: String) {
module
.register_async_method(
Box::leak(method_name.into_boxed_str()),
move |params, ctx, extensions| {
let handler = self.clone();
async move {
let ctx =
match Self::Ctx::from_request_extensions((*ctx).clone(), extensions)
.await
{
Ok(ctx) => ctx,
Err(e) => {
return ResponsePayload::error(e);
}
};
let result = handler.call(ctx, params);
ResponsePayload::success(result.transform())
}
},
)
.unwrap();
}
}
impl<Ctx, T, MSig, MValue>
RegisterableHandler<Ctx, MSig, MValue, marker::MFuture<marker::MResponse<MValue>>> for T
where
Ctx: 'static + Clone + Send + Sync,
MValue: marker::ResponseMarker,
T: QubitHandler<Ctx, MSig>,
T::Return: Future + Send,
<T::Return as Future>::Output: ResponseValue<MValue>,
{
type Response = <T::Return as Future>::Output;
fn register(self, module: &mut RpcModule<Ctx>, method_name: String) {
module
.register_async_method(
Box::leak(method_name.into_boxed_str()),
move |params, ctx, extensions| {
let f = self.clone();
async move {
let ctx =
match Self::Ctx::from_request_extensions((*ctx).clone(), extensions)
.await
{
Ok(ctx) => ctx,
Err(e) => {
return ResponsePayload::error(e);
}
};
let result = f.call(ctx, params).await;
ResponsePayload::success(result.transform())
}
},
)
.unwrap();
}
}
impl<Ctx, T, MValue, MSig> RegisterableHandler<Ctx, MSig, MValue, marker::MStream<MValue>> for T
where
Ctx: 'static + Clone + Send + Sync,
MValue: marker::ResponseMarker,
T: QubitHandler<Ctx, MSig>,
T::Return: Stream + Send,
<T::Return as Stream>::Item: Send + ResponseValue<MValue>,
{
type Response = <T::Return as Stream>::Item;
fn register(self, module: &mut RpcModule<Ctx>, method_name: String) {
let notif_method_name = format!("{method_name}_notif");
let unsub_method_name = format!("{method_name}_unsub");
module
.register_subscription(
Box::leak(method_name.into_boxed_str()),
Box::leak(notif_method_name.into_boxed_str()),
Box::leak(unsub_method_name.into_boxed_str()),
move |params, pending, ctx, extensions| {
let f = self.clone();
async move {
let ctx =
match Self::Ctx::from_request_extensions((*ctx).clone(), extensions)
.await
{
Ok(ctx) => ctx,
Err(e) => {
pending.reject(e).await;
return SubscriptionCloseResponse::None;
}
};
let sink = pending.accept().await.unwrap();
let mut count = 0;
let subscription_id = sink.subscription_id();
let mut stream = pin!(f.call(ctx, params));
while let Some(item) = stream.next().await {
let item = serde_json::value::to_raw_value(&item.transform()).unwrap();
if let Some(DisconnectError(..)) = sink.send(item).await.err() {
break;
};
count += 1;
}
SubscriptionCloseResponse::Notif(SubscriptionMessage::from(
serde_json::value::to_raw_value(
&json!({ "close_stream": subscription_id, "count": count }),
)
.unwrap(),
))
}
},
)
.unwrap();
}
}
impl<Ctx, T, MValue, MSig>
RegisterableHandler<Ctx, MSig, MValue, marker::MFuture<marker::MStream<MValue>>> for T
where
Ctx: 'static + Clone + Send + Sync,
MValue: marker::ResponseMarker,
T: QubitHandler<Ctx, MSig>,
T::Return: Send + Future,
<T::Return as Future>::Output: Stream + Send,
<<T::Return as Future>::Output as Stream>::Item: Send + ResponseValue<MValue>,
{
type Response = <<T::Return as Future>::Output as Stream>::Item;
fn register(self, module: &mut RpcModule<Ctx>, method_name: String) {
let notif_method_name = format!("{method_name}_notif");
let unsub_method_name = format!("{method_name}_unsub");
module
.register_subscription(
Box::leak(method_name.into_boxed_str()),
Box::leak(notif_method_name.into_boxed_str()),
Box::leak(unsub_method_name.into_boxed_str()),
move |params, pending, ctx, extensions| {
let f = self.clone();
async move {
let ctx =
match Self::Ctx::from_request_extensions((*ctx).clone(), extensions)
.await
{
Ok(ctx) => ctx,
Err(e) => {
pending.reject(e).await;
return SubscriptionCloseResponse::None;
}
};
let sink = pending.accept().await.unwrap();
let mut count = 0;
let subscription_id = sink.subscription_id();
let mut stream = pin!(f.call(ctx, params).await);
while let Some(item) = stream.next().await {
let item = serde_json::value::to_raw_value(&item.transform()).unwrap();
if let Some(DisconnectError(..)) = sink.send(item).await.err() {
break;
};
count += 1;
}
SubscriptionCloseResponse::Notif(SubscriptionMessage::from(
serde_json::value::to_raw_value(
&json!({ "close_stream": subscription_id, "count": count }),
)
.unwrap(),
))
}
},
)
.unwrap();
}
}
#[cfg(test)]
mod test {
use crate::{RpcError, handler::ts::TypeCollector};
use super::{ctx::FromRequestExtensions, *};
use futures::stream;
use rstest::rstest;
use serde_json::{Value, json};
use std::{fmt::Debug, iter};
mod register {
use jsonrpsee::RpcModule;
use serde::Deserialize;
use super::*;
fn simple_iter() -> impl Iterator<Item = usize> {
0..3
}
fn register_handler<
F,
MSig,
MValue: marker::ResponseMarker,
MReturn: marker::HandlerReturnMarker,
>(
handler: F,
) -> RpcModule<()>
where
F: RegisterableHandler<(), MSig, MValue, MReturn, Ctx = ()>,
{
let mut module = RpcModule::new(());
F::register(handler, &mut module, "handler".to_string());
module
}
async fn test_handler<
F,
MSig,
MValue: marker::ResponseMarker,
MReturn: marker::HandlerReturnMarker,
>(
handler: F,
) -> <F::Response as ResponseValue<MValue>>::Value
where
F: RegisterableHandler<(), MSig, MValue, MReturn, Ctx = ()>,
<F::Response as ResponseValue<MValue>>::Value: for<'a> Deserialize<'a>,
{
let module = register_handler(handler);
let fut = module
.call::<[(); 0], <F::Response as ResponseValue<MValue>>::Value>("handler", []);
fut.await.unwrap()
}
#[tokio::test]
async fn ts() {
assert_eq!(test_handler(|| 123u32).await, 123);
}
#[tokio::test]
async fn iter() {
assert_eq!(test_handler(simple_iter).await, vec![0, 1, 2]);
}
#[tokio::test]
async fn stream() {
let module = register_handler(|| futures::stream::iter(simple_iter()));
let mut subs = module.subscribe("handler", [] as [(); 0], 3).await.unwrap();
let mut next = async || subs.next::<usize>().await.unwrap().unwrap().0;
assert_eq!(0, next().await);
assert_eq!(1, next().await);
assert_eq!(2, next().await);
assert_eq!(
subs.next::<Value>().await.unwrap().unwrap().0["count"]
.as_i64()
.unwrap(),
3
);
assert!(subs.next::<Value>().await.is_none());
}
}
#[rstest]
#[case::ts_value(|| 123)]
#[case::async_ts_value(|| async { 123 })]
#[case::stream(|| stream::once(async { 123 }))]
#[case::async_stream(|| async { stream::once(async { 123 }) })]
#[case::iter(|| iter::once(123))]
#[case::async_iter(|| async { iter::once(123) })]
#[case::stream_iter(|| stream::once(async { iter::once(123) }))]
#[case::async_stream_iter(|| async { stream::once(async { iter::once(123) }) })]
#[case::iter_iter(|| iter::once(iter::once(123)))]
#[case::async_iter_iter(|| async { iter::once(iter::once(123)) })]
#[case::stream_iter_iter(|| stream::once(async { iter::once(iter::once(123)) }))]
#[case::async_stream_iter_iter(|| async { stream::once(async { iter::once(iter::once(123)) }) })]
fn register_handler<
MSig,
MValue: marker::ResponseMarker,
MReturn: marker::HandlerReturnMarker,
>(
#[case] handler: impl RegisterableHandler<(), MSig, MValue, MReturn, Ctx = ()>,
) {
handler.register(&mut RpcModule::new(()), "handler".to_string());
}
#[rstest]
#[case(|| {}, json!([]), ())]
#[case(|_ctx: ()| {}, json!([]), ())]
#[case(|_ctx: (), param: u32| param, json!([123]), 123)]
#[case(|_ctx: (), param_1: u32, param_2: String| -> (u32, String) { (param_1, param_2) }, json!([123, "hello"]), (123, "hello".to_string()))]
fn call_handler<H, MSig>(#[case] handler: H, #[case] params: Value, #[case] expected: H::Return)
where
H: QubitHandler<(), MSig, Ctx = ()>,
H::Return: Debug + PartialEq,
{
let output = handler.call(
(),
Params::new(Some(&serde_json::to_string(¶ms).unwrap())).into_owned(),
);
assert_eq!(output, expected);
}
#[derive(Clone)]
struct SampleCtx;
#[derive(Clone)]
struct DerivedCtx;
impl FromRequestExtensions<SampleCtx> for DerivedCtx {
async fn from_request_extensions(
_ctx: SampleCtx,
_extensions: http::Extensions,
) -> Result<Self, RpcError> {
Ok(DerivedCtx)
}
}
#[test]
fn derived_ctx() {
fn handler(_ctx: DerivedCtx) {}
handler.register(&mut RpcModule::new(SampleCtx), "handler".to_string());
}
#[rstest]
#[case::unit_handler(|| {}, (), [], "null")]
#[case::unit_handler_other_ctx(|| {}, SampleCtx, [], "null")]
#[case::single_ctx_param(|_ctx: SampleCtx| {}, SampleCtx, [], "null")]
#[case::only_return_ty(|| -> bool { todo!() }, (), [], "boolean")]
#[case::ctx_and_param(|_ctx: SampleCtx, _a: u32| {}, SampleCtx, ["number"], "null")]
#[case::ctx_and_param_and_return(|_ctx: SampleCtx, _a: u32| -> bool { todo!() }, SampleCtx, ["number"], "boolean")]
#[case::ctx_and_multi_param(|_ctx: SampleCtx, _a: u32, _b: String, _c: bool| {}, SampleCtx, ["number", "string", "boolean"], "null")]
#[case::ctx_and_multi_param_return(|_ctx: SampleCtx, _a: u32, _b: String, _c: bool| -> bool { todo!() }, SampleCtx, ["number", "string", "boolean"], "boolean")]
#[case::produce_iter(|| { [1, 2, 3].into_iter() }, (), [], "Array<number>")]
#[case::produce_stream(|| { stream::iter([1, 2, 3]) }, (), [], "number")]
fn handler_ts_type<H, Ctx, MSig, MValue, MReturn>(
#[case] _handler: H,
#[case] _ctx: Ctx,
#[case] expected_params: impl IntoIterator<Item = &'static str>,
#[case] expected_return: &'static str,
) where
MValue: marker::ResponseMarker,
MReturn: marker::HandlerReturnMarker,
H: RegisterableHandler<Ctx, MSig, MValue, MReturn>,
Ctx: 'static + Clone + Send + Sync,
H::Ctx: 'static + Send + Sync + FromRequestExtensions<Ctx>,
{
assert_eq!(
TypeCollector::collect_names::<H::Params>(),
expected_params.into_iter().collect::<Vec<_>>()
);
assert_eq!(
<<H::Response as ResponseValue<_>>::Value as TS>::name(&ts_rs::Config::default()),
expected_return
);
}
mod handler_traits {
use super::*;
use static_assertions::assert_impl_all;
assert_impl_all!(
fn () -> (): QubitHandler<(), ((), ()), Ctx = (), Params = (), Return = ()>
);
assert_impl_all!(
fn (u32) -> (): QubitHandler<u32, ((u32,), ()), Ctx = u32, Params = (), Return = ()>
);
assert_impl_all!(
fn (u32, String, bool) -> (): QubitHandler<u32, ((u32, String, bool), ()), Ctx = u32, Params = (String, bool), Return = ()>
);
assert_impl_all!(
fn () -> u32: QubitHandler<(), ((), u32), Ctx = (), Params = (), Return = u32>
);
assert_impl_all!(
fn () -> std::vec::IntoIter<u32> : QubitHandler<(), ((), std::vec::IntoIter<u32>)>
);
assert_impl_all!(
fn () -> futures::stream::Iter<std::vec::IntoIter<u32>> : QubitHandler<(), ((), futures::stream::Iter<std::vec::IntoIter<u32>>)>
);
assert_impl_all!(
fn () -> futures::stream::Iter<std::vec::IntoIter<std::vec::IntoIter<u32>>> : QubitHandler<(), ((), futures::stream::Iter<std::vec::IntoIter<std::vec::IntoIter<u32>>>)>
);
}
}