use std::{
fmt::{self, Debug, Formatter},
future::Future,
};
use crate::{
body::Body,
decode,
request::{self, BoxRequest},
response::BoxResponse,
Response,
};
use self::{
boxed::BoxedTransport,
transport::{SocketChannels, SocketRequestMarker, TransportError},
};
use super::Request;
use error::*;
use futures_util::TryFutureExt;
use socket::*;
use tower::{Layer, Service};
pub mod boxed;
pub mod error;
pub mod layer;
pub mod socket;
pub mod transport;
#[doc(hidden)]
pub mod prelude {
pub use super::{
error::{ClientError, ClientResult},
socket::Socket,
transport::TransportError,
Client,
};
pub use crate::{
request::{BoxRequest, IntoRequest, Request},
response::{BoxResponse, Response},
};
pub use std::{borrow::Cow, convert::TryInto, fmt::Debug, future::Future};
pub use tower::Service;
}
pub struct Client<Inner> {
transport: Inner,
}
impl<Inner: Debug> Debug for Client<Inner> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("Client")
.field("inner", &self.transport)
.finish()
}
}
impl<Inner: Clone> Clone for Client<Inner> {
fn clone(&self) -> Self {
Self {
transport: self.transport.clone(),
}
}
}
impl<Inner> Client<Inner> {
pub fn new(transport: Inner) -> Client<Inner> {
Client { transport }
}
}
impl<Inner, InnerErr> Client<Inner>
where
Inner: Service<BoxRequest, Response = BoxResponse, Error = TransportError<InnerErr>>
+ Send
+ Clone
+ 'static,
Inner::Future: Send,
InnerErr: std::error::Error + Sync + Send + 'static,
{
pub fn boxed(self) -> Client<BoxedTransport> {
Client {
transport: BoxedTransport::new(self.transport),
}
}
}
impl<Inner, InnerErr> Client<Inner>
where
Inner: Service<BoxRequest, Response = BoxResponse, Error = TransportError<InnerErr>> + 'static,
InnerErr: 'static,
{
pub fn layer<S, L>(self, l: L) -> Client<S>
where
L: Layer<Inner, Service = S>,
S: Service<BoxRequest>,
{
Client {
transport: l.layer(self.transport),
}
}
pub fn execute_request<Req, Resp>(
&mut self,
req: Request<Req>,
) -> impl Future<Output = ClientResult<Response<Resp>, InnerErr>> + 'static
where
Req: prost::Message,
Resp: prost::Message + Default,
{
Service::call(&mut self.transport, req.map::<()>())
.map_ok(|resp| resp.map::<Resp>())
.map_err(ClientError::from)
}
pub fn connect_socket<Req, Resp>(
&mut self,
mut req: Request<()>,
) -> impl Future<Output = ClientResult<Socket<Req, Resp>, InnerErr>> + 'static
where
Req: prost::Message,
Resp: prost::Message + Default,
{
req.extensions_mut().insert(SocketRequestMarker);
Service::call(&mut self.transport, req)
.map_ok(|mut resp| {
let chans = resp
.extensions_mut()
.remove::<SocketChannels>()
.expect("transport did not return socket channels - this is a bug");
Socket::new(
chans.rx,
chans.tx,
socket::encode_message,
socket::decode_message,
)
})
.map_err(ClientError::from)
}
pub fn connect_socket_req<Req, Resp>(
&mut self,
request: Request<Req>,
) -> impl Future<Output = ClientResult<Socket<Req, Resp>, InnerErr>> + 'static
where
Req: prost::Message + Default + 'static,
Resp: prost::Message + Default + 'static,
{
let request::Parts {
body,
extensions,
endpoint,
..
} = request.into();
let request: BoxRequest = Request::from(request::Parts {
body: Body::empty(),
endpoint: endpoint.clone(),
extensions,
});
let connect_fut = self.connect_socket(request);
async move {
let mut socket = connect_fut.await?;
let message = decode::decode_body(body).await?;
socket
.send_message(message)
.await
.map_err(|err| match err {
SocketError::MessageDecode(err) => ClientError::MessageDecode(err),
SocketError::Protocol(err) => ClientError::EndpointError {
hrpc_error: err,
endpoint,
},
SocketError::Transport(err) => ClientError::EndpointError {
hrpc_error: HrpcError::from(err).with_identifier("hrpcrs.socket-error"),
endpoint,
},
})?;
Ok(socket)
}
}
}