use cratestack_codec_cbor::CborCodec;
use cratestack_core::rpc::{RPC_BATCH_PATH, RpcRequest, RpcResponseFrame};
use reqwest::Method;
use serde::Serialize;
use serde::de::DeserializeOwned;
use crate::client::CratestackClient;
use crate::codec::HttpClientCodec;
use crate::config::ClientConfig;
use crate::rpc::batch::BatchBuilder;
use crate::rpc::error::{RpcClientError, client_error_to_rpc, decode_rpc_unary_response};
use crate::streaming::pump_streamed_response_typed;
const RPC_BATCH_PATH_PLAIN: &str = RPC_BATCH_PATH;
#[derive(Clone)]
pub struct RpcClient<C = CborCodec> {
pub(crate) inner: CratestackClient<C>,
}
impl RpcClient<CborCodec> {
pub fn cbor(config: ClientConfig) -> Self {
Self::new(CratestackClient::cbor(config))
}
}
impl<C> RpcClient<C>
where
C: HttpClientCodec + Clone,
{
pub fn new(inner: CratestackClient<C>) -> Self {
Self { inner }
}
pub fn inner(&self) -> &CratestackClient<C> {
&self.inner
}
pub fn batch_builder(&self) -> BatchBuilder<C> {
BatchBuilder::new(self.clone())
}
pub async fn call<I, O>(&self, op_id: &str, input: &I) -> Result<O, RpcClientError>
where
I: Serialize,
O: DeserializeOwned,
{
let body = self
.inner
.codec
.encode(input)
.map_err(RpcClientError::Codec)?;
let path = format!("/rpc/{}", op_id);
let response = self
.inner
.request_raw_with_query_and_accept(Method::POST, &path, Some(body), None, &[], None)
.await
.map_err(client_error_to_rpc)?;
decode_rpc_unary_response(&self.inner.codec, &response)
}
pub async fn batch(
&self,
requests: &[RpcRequest],
) -> Result<Vec<RpcResponseFrame>, RpcClientError> {
let body = self
.inner
.codec
.encode(&requests)
.map_err(RpcClientError::Codec)?;
let response = self
.inner
.request_raw_with_query_and_accept(
Method::POST,
RPC_BATCH_PATH_PLAIN,
Some(body),
None,
&[],
None,
)
.await
.map_err(client_error_to_rpc)?;
decode_rpc_unary_response::<C, Vec<RpcResponseFrame>>(&self.inner.codec, &response)
}
pub async fn call_streaming<I, O>(
&self,
op_id: &str,
input: &I,
) -> Result<tokio::sync::mpsc::Receiver<Result<O, RpcClientError>>, RpcClientError>
where
I: Serialize,
O: DeserializeOwned + Send + 'static,
{
let body = self
.inner
.codec
.encode(input)
.map_err(RpcClientError::Codec)?;
let path = format!("/rpc/{}", op_id);
let response = self
.inner
.request_streamed_with_query_and_accept(
Method::POST,
&path,
Some(body),
None,
&[],
self.inner.codec.sequence_accept_header_value(),
)
.await
.map_err(client_error_to_rpc)?;
let (tx, rx) = tokio::sync::mpsc::channel(16);
tokio::spawn(pump_streamed_response_typed::<O, RpcClientError, _>(
response,
tx,
client_error_to_rpc,
));
Ok(rx)
}
}