use std::{collections::BTreeSet, future::Future, marker::PhantomData};
use prost::Message;
use super::MethodDescriptor;
pub trait Rpc {
type Request: Message;
type Response: Message + Default;
const METHOD: &'static MethodDescriptor;
}
pub trait UnaryRpc: Rpc {}
pub trait ServerStreamingRpc: Rpc {}
pub trait ClientStreamingRpc: Rpc {}
pub trait BidirectionalRpc: Rpc {}
pub trait MessageReader: Send {
type Error: std::error::Error + Send + Sync + 'static;
fn next(&mut self) -> impl Future<Output = Result<Option<Vec<u8>>, Self::Error>> + Send;
fn cancel(&mut self);
}
pub trait MessageWriter: Send {
type Error: std::error::Error + Send + Sync + 'static;
fn send(&mut self, message: Vec<u8>) -> impl Future<Output = Result<(), Self::Error>> + Send;
fn finish(&mut self) -> impl Future<Output = Result<(), Self::Error>> + Send;
fn abort(&mut self);
}
pub trait RpcTransport: Send + Sync {
type Error: std::error::Error + Send + Sync + 'static;
type Reader: MessageReader<Error = Self::Error>;
type Writer: MessageWriter<Error = Self::Error>;
fn unary(
&self,
method: &'static MethodDescriptor,
request: Vec<u8>,
) -> impl Future<Output = Result<Vec<u8>, Self::Error>> + Send;
fn observe(
&self,
method: &'static MethodDescriptor,
request: Vec<u8>,
) -> impl Future<Output = Result<Self::Reader, Self::Error>> + Send;
fn exchange(
&self,
method: &'static MethodDescriptor,
opening: Vec<u8>,
) -> impl Future<Output = Result<(Self::Writer, Self::Reader), Self::Error>> + Send;
}
#[derive(Debug, thiserror::Error)]
pub enum ClientError<E: std::error::Error> {
#[error("endpoint does not implement {0}")]
NotImplemented(&'static str),
#[error("incompatible peer for {0}")]
Protocol(&'static str),
#[error("{0} requires a stable client operation ID")]
MissingOperationId(&'static str),
#[error("invalid request metadata: {0}")]
Metadata(#[from] crate::RequestMetadataError),
#[error("transport failure: {0}")]
Transport(E),
#[error("invalid protobuf response: {0}")]
Decode(#[from] prost::DecodeError),
}
pub struct Client<T> {
transport: T,
implemented: BTreeSet<String>,
protocol: Option<crate::heddle::api::common::ProtocolCompatibility>,
}
impl<T: RpcTransport> Client<T> {
pub fn new(transport: T, implemented: impl IntoIterator<Item = String>) -> Self {
Self {
transport,
implemented: implemented.into_iter().collect(),
protocol: None,
}
}
pub fn with_protocol(
mut self,
protocol: crate::heddle::api::common::ProtocolCompatibility,
) -> Self {
self.protocol = Some(protocol);
self
}
fn encode<M: Rpc>(&self, request: &M::Request) -> Result<Vec<u8>, ClientError<T::Error>> {
let method = M::METHOD;
if !self.implemented.contains(method.path) {
return Err(ClientError::NotImplemented(method.path));
}
if !method.mandatory_features.is_empty() {
crate::import_authority::require_hybrid_peer(self.protocol.as_ref())
.map_err(|_| ClientError::Protocol(method.path))?;
}
let bytes = request.encode_to_vec();
validate_stream_protocol(method, &bytes, true, true)
.map_err(|_| ClientError::Protocol(method.path))?;
if method.client_operation_id_required {
let Some(field) = method.client_operation_id_field_number else {
return Err(ClientError::MissingOperationId(method.path));
};
let id = crate::transport::protobuf_string_field(&bytes, field)?;
if id.is_none_or(|value| value.trim().is_empty()) {
return Err(ClientError::MissingOperationId(method.path));
}
}
Ok(bytes)
}
pub async fn call<M: UnaryRpc>(
&self,
request: &M::Request,
) -> Result<M::Response, ClientError<T::Error>> {
let bytes = self
.transport
.unary(M::METHOD, self.encode::<M>(request)?)
.await
.map_err(ClientError::Transport)?;
Ok(M::Response::decode(bytes.as_slice())?)
}
pub async fn observe<M: ServerStreamingRpc>(
&self,
request: &M::Request,
) -> Result<Messages<T::Reader, M::Response>, ClientError<T::Error>> {
let reader = self
.transport
.observe(M::METHOD, self.encode::<M>(request)?)
.await
.map_err(ClientError::Transport)?;
Ok(Messages {
reader,
method: M::METHOD,
first: true,
done: false,
message: PhantomData,
})
}
pub async fn exchange<M: BidirectionalRpc>(
&self,
opening: &M::Request,
) -> Result<
(
Sender<T::Writer, M::Request>,
Messages<T::Reader, M::Response>,
),
ClientError<T::Error>,
> {
let (writer, reader) = self
.transport
.exchange(M::METHOD, self.encode::<M>(opening)?)
.await
.map_err(ClientError::Transport)?;
Ok((
Sender {
writer,
method: M::METHOD,
finished: false,
message: PhantomData,
},
Messages {
reader,
method: M::METHOD,
first: true,
done: false,
message: PhantomData,
},
))
}
}
pub struct Messages<R: MessageReader, O> {
reader: R,
method: &'static MethodDescriptor,
first: bool,
done: bool,
message: PhantomData<O>,
}
impl<R: MessageReader, O> Messages<R, O> {
pub fn cancel(&mut self) {
if !self.done {
self.reader.cancel();
self.done = true;
}
}
}
impl<R: MessageReader, O: Message + Default> Messages<R, O> {
pub async fn next(&mut self) -> Result<Option<O>, ClientError<R::Error>> {
if self.done {
return Ok(None);
}
let decoded = match self.reader.next().await {
Ok(Some(bytes)) => {
if validate_stream_protocol(self.method, &bytes, false, self.first).is_err() {
Err(ClientError::Protocol(self.method.path))
} else {
self.first = false;
O::decode(bytes.as_slice())
.map(Some)
.map_err(ClientError::Decode)
}
}
Ok(None) => {
if self.first && is_hybrid_stream(self.method) {
Err(ClientError::Protocol(self.method.path))
} else {
self.done = true;
Ok(None)
}
}
Err(error) => Err(ClientError::Transport(error)),
};
if decoded.is_err() {
self.reader.cancel();
self.done = true;
}
decoded
}
}
impl<R: MessageReader, O> Drop for Messages<R, O> {
fn drop(&mut self) {
self.cancel();
}
}
pub struct Sender<W: MessageWriter, I> {
writer: W,
method: &'static MethodDescriptor,
finished: bool,
message: PhantomData<I>,
}
impl<W: MessageWriter, I: Message> Sender<W, I> {
pub async fn send(&mut self, message: &I) -> Result<(), ClientError<W::Error>> {
let bytes = message.encode_to_vec();
validate_stream_protocol(self.method, &bytes, true, false)
.map_err(|_| ClientError::Protocol(self.method.path))?;
self.writer
.send(bytes)
.await
.map_err(ClientError::Transport)
}
pub async fn finish(mut self) -> Result<(), W::Error> {
self.writer.finish().await?;
self.finished = true;
Ok(())
}
}
impl<W: MessageWriter, I> Drop for Sender<W, I> {
fn drop(&mut self) {
if !self.finished {
self.writer.abort();
}
}
}
fn is_hybrid_stream(method: &MethodDescriptor) -> bool {
method.path.starts_with("/heddle.api.v1alpha2.SyncService/")
&& !method.mandatory_features.is_empty()
}
fn validate_stream_protocol(
method: &MethodDescriptor,
bytes: &[u8],
request: bool,
first: bool,
) -> Result<(), crate::hybrid_codec::Reject> {
use crate::heddle::api::v1alpha2::*;
use crate::hybrid_codec::Reject;
use crate::import_authority::require_hybrid_peer;
if !is_hybrid_stream(method) {
return Ok(());
}
macro_rules! frame {
($ty:ty, $variant:path) => {{
let value = <$ty>::decode(bytes).map_err(|_| Reject::Protocol)?;
match value.body {
Some($variant(open)) => require_hybrid_peer(open.protocol.as_ref()),
_ if first => Err(Reject::Protocol),
_ => Ok(()),
}
}};
}
match (method.path.rsplit('/').next(), request) {
(Some("Fetch"), true) => {
frame!(FetchClientFrame, fetch_client_frame::Body::Open)
}
(Some("Fetch"), false) => frame!(FetchServerFrame, fetch_server_frame::Body::Ready),
(Some("PublishContent"), true) => frame!(
PublishContentClientFrame,
publish_content_client_frame::Body::Open
),
(Some("PublishContent"), false) => frame!(
PublishContentServerFrame,
publish_content_server_frame::Body::Ready
),
(Some("ReplicateThread"), true) => {
frame!(ReplicateThreadRequest, replicate_thread_request::Body::Open)
}
(Some("ReplicateThread"), false) => frame!(
ReplicateThreadResponse,
replicate_thread_response::Body::Ready
),
_ => Err(Reject::Protocol),
}
}