1use std::{collections::BTreeSet, future::Future, marker::PhantomData};
3
4use prost::Message;
5
6use super::MethodDescriptor;
7
8pub trait Rpc {
9 type Request: Message;
10 type Response: Message + Default;
11 const METHOD: &'static MethodDescriptor;
12}
13pub trait UnaryRpc: Rpc {}
14pub trait ServerStreamingRpc: Rpc {}
15pub trait ClientStreamingRpc: Rpc {}
16pub trait BidirectionalRpc: Rpc {}
17
18pub trait MessageReader: Send {
21 type Error: std::error::Error + Send + Sync + 'static;
22 fn next(&mut self) -> impl Future<Output = Result<Option<Vec<u8>>, Self::Error>> + Send;
23 fn cancel(&mut self);
25}
26
27pub trait MessageWriter: Send {
28 type Error: std::error::Error + Send + Sync + 'static;
29 fn send(&mut self, message: Vec<u8>) -> impl Future<Output = Result<(), Self::Error>> + Send;
30 fn finish(&mut self) -> impl Future<Output = Result<(), Self::Error>> + Send;
31 fn abort(&mut self);
32}
33
34pub trait RpcTransport: Send + Sync {
38 type Error: std::error::Error + Send + Sync + 'static;
39 type Reader: MessageReader<Error = Self::Error>;
40 type Writer: MessageWriter<Error = Self::Error>;
41 fn unary(
42 &self,
43 method: &'static MethodDescriptor,
44 request: Vec<u8>,
45 ) -> impl Future<Output = Result<Vec<u8>, Self::Error>> + Send;
46 fn observe(
47 &self,
48 method: &'static MethodDescriptor,
49 request: Vec<u8>,
50 ) -> impl Future<Output = Result<Self::Reader, Self::Error>> + Send;
51 fn exchange(
52 &self,
53 method: &'static MethodDescriptor,
54 opening: Vec<u8>,
55 ) -> impl Future<Output = Result<(Self::Writer, Self::Reader), Self::Error>> + Send;
56}
57
58#[derive(Debug, thiserror::Error)]
59pub enum ClientError<E: std::error::Error> {
60 #[error("endpoint does not implement {0}")]
61 NotImplemented(&'static str),
62 #[error("incompatible peer for {0}")]
63 Protocol(&'static str),
64 #[error("{0} requires a stable client operation ID")]
65 MissingOperationId(&'static str),
66 #[error("invalid request metadata: {0}")]
67 Metadata(#[from] crate::RequestMetadataError),
68 #[error("transport failure: {0}")]
69 Transport(E),
70 #[error("invalid protobuf response: {0}")]
71 Decode(#[from] prost::DecodeError),
72}
73
74pub struct Client<T> {
75 transport: T,
76 implemented: BTreeSet<String>,
77 protocol: Option<crate::heddle::api::common::ProtocolCompatibility>,
78}
79
80impl<T: RpcTransport> Client<T> {
81 pub fn new(transport: T, implemented: impl IntoIterator<Item = String>) -> Self {
84 Self {
85 transport,
86 implemented: implemented.into_iter().collect(),
87 protocol: None,
88 }
89 }
90
91 pub fn with_protocol(
95 mut self,
96 protocol: crate::heddle::api::common::ProtocolCompatibility,
97 ) -> Self {
98 self.protocol = Some(protocol);
99 self
100 }
101
102 fn encode<M: Rpc>(&self, request: &M::Request) -> Result<Vec<u8>, ClientError<T::Error>> {
103 let method = M::METHOD;
104 if !self.implemented.contains(method.path) {
105 return Err(ClientError::NotImplemented(method.path));
106 }
107 if !method.mandatory_features.is_empty() {
108 crate::import_authority::require_hybrid_peer(self.protocol.as_ref())
109 .map_err(|_| ClientError::Protocol(method.path))?;
110 }
111 let bytes = request.encode_to_vec();
112 validate_stream_protocol(method, &bytes, true, true)
113 .map_err(|_| ClientError::Protocol(method.path))?;
114 if method.client_operation_id_required {
115 let Some(field) = method.client_operation_id_field_number else {
116 return Err(ClientError::MissingOperationId(method.path));
117 };
118 let id = crate::transport::protobuf_string_field(&bytes, field)?;
119 if id.is_none_or(|value| value.trim().is_empty()) {
120 return Err(ClientError::MissingOperationId(method.path));
121 }
122 }
123 Ok(bytes)
124 }
125
126 pub async fn call<M: UnaryRpc>(
127 &self,
128 request: &M::Request,
129 ) -> Result<M::Response, ClientError<T::Error>> {
130 let bytes = self
131 .transport
132 .unary(M::METHOD, self.encode::<M>(request)?)
133 .await
134 .map_err(ClientError::Transport)?;
135 Ok(M::Response::decode(bytes.as_slice())?)
136 }
137
138 pub async fn observe<M: ServerStreamingRpc>(
139 &self,
140 request: &M::Request,
141 ) -> Result<Messages<T::Reader, M::Response>, ClientError<T::Error>> {
142 let reader = self
143 .transport
144 .observe(M::METHOD, self.encode::<M>(request)?)
145 .await
146 .map_err(ClientError::Transport)?;
147 Ok(Messages {
148 reader,
149 method: M::METHOD,
150 first: true,
151 done: false,
152 message: PhantomData,
153 })
154 }
155
156 pub async fn exchange<M: BidirectionalRpc>(
158 &self,
159 opening: &M::Request,
160 ) -> Result<
161 (
162 Sender<T::Writer, M::Request>,
163 Messages<T::Reader, M::Response>,
164 ),
165 ClientError<T::Error>,
166 > {
167 let (writer, reader) = self
168 .transport
169 .exchange(M::METHOD, self.encode::<M>(opening)?)
170 .await
171 .map_err(ClientError::Transport)?;
172 Ok((
173 Sender {
174 writer,
175 method: M::METHOD,
176 finished: false,
177 message: PhantomData,
178 },
179 Messages {
180 reader,
181 method: M::METHOD,
182 first: true,
183 done: false,
184 message: PhantomData,
185 },
186 ))
187 }
188}
189
190pub struct Messages<R: MessageReader, O> {
191 reader: R,
192 method: &'static MethodDescriptor,
193 first: bool,
194 done: bool,
195 message: PhantomData<O>,
196}
197
198impl<R: MessageReader, O> Messages<R, O> {
199 pub fn cancel(&mut self) {
202 if !self.done {
203 self.reader.cancel();
204 self.done = true;
205 }
206 }
207}
208
209impl<R: MessageReader, O: Message + Default> Messages<R, O> {
210 pub async fn next(&mut self) -> Result<Option<O>, ClientError<R::Error>> {
211 if self.done {
212 return Ok(None);
213 }
214 let decoded = match self.reader.next().await {
215 Ok(Some(bytes)) => {
216 if validate_stream_protocol(self.method, &bytes, false, self.first).is_err() {
217 Err(ClientError::Protocol(self.method.path))
218 } else {
219 self.first = false;
220 O::decode(bytes.as_slice())
221 .map(Some)
222 .map_err(ClientError::Decode)
223 }
224 }
225 Ok(None) => {
226 if self.first && is_hybrid_stream(self.method) {
227 Err(ClientError::Protocol(self.method.path))
228 } else {
229 self.done = true;
230 Ok(None)
231 }
232 }
233 Err(error) => Err(ClientError::Transport(error)),
234 };
235 if decoded.is_err() {
236 self.reader.cancel();
237 self.done = true;
238 }
239 decoded
240 }
241}
242
243impl<R: MessageReader, O> Drop for Messages<R, O> {
244 fn drop(&mut self) {
245 self.cancel();
246 }
247}
248
249pub struct Sender<W: MessageWriter, I> {
250 writer: W,
251 method: &'static MethodDescriptor,
252 finished: bool,
253 message: PhantomData<I>,
254}
255
256impl<W: MessageWriter, I: Message> Sender<W, I> {
257 pub async fn send(&mut self, message: &I) -> Result<(), ClientError<W::Error>> {
258 let bytes = message.encode_to_vec();
259 validate_stream_protocol(self.method, &bytes, true, false)
260 .map_err(|_| ClientError::Protocol(self.method.path))?;
261 self.writer
262 .send(bytes)
263 .await
264 .map_err(ClientError::Transport)
265 }
266 pub async fn finish(mut self) -> Result<(), W::Error> {
267 self.writer.finish().await?;
268 self.finished = true;
269 Ok(())
270 }
271}
272
273impl<W: MessageWriter, I> Drop for Sender<W, I> {
274 fn drop(&mut self) {
275 if !self.finished {
276 self.writer.abort();
277 }
278 }
279}
280
281fn is_hybrid_stream(method: &MethodDescriptor) -> bool {
284 method.path.starts_with("/heddle.api.v1alpha2.SyncService/")
285 && !method.mandatory_features.is_empty()
286}
287fn validate_stream_protocol(
289 method: &MethodDescriptor,
290 bytes: &[u8],
291 request: bool,
292 first: bool,
293) -> Result<(), crate::hybrid_codec::Reject> {
294 use crate::heddle::api::v1alpha2::*;
295 use crate::hybrid_codec::Reject;
296 use crate::import_authority::require_hybrid_peer;
297 if !is_hybrid_stream(method) {
298 return Ok(());
299 }
300 macro_rules! frame {
301 ($ty:ty, $variant:path) => {{
302 let value = <$ty>::decode(bytes).map_err(|_| Reject::Protocol)?;
303 match value.body {
304 Some($variant(open)) => require_hybrid_peer(open.protocol.as_ref()),
305 _ if first => Err(Reject::Protocol),
306 _ => Ok(()),
307 }
308 }};
309 }
310 match (method.path.rsplit('/').next(), request) {
311 (Some("Fetch"), true) => {
312 frame!(FetchClientFrame, fetch_client_frame::Body::Open)
313 }
314 (Some("Fetch"), false) => frame!(FetchServerFrame, fetch_server_frame::Body::Ready),
315 (Some("PublishContent"), true) => frame!(
316 PublishContentClientFrame,
317 publish_content_client_frame::Body::Open
318 ),
319 (Some("PublishContent"), false) => frame!(
320 PublishContentServerFrame,
321 publish_content_server_frame::Body::Ready
322 ),
323 (Some("ReplicateThread"), true) => {
324 frame!(ReplicateThreadRequest, replicate_thread_request::Body::Open)
325 }
326 (Some("ReplicateThread"), false) => frame!(
327 ReplicateThreadResponse,
328 replicate_thread_response::Body::Ready
329 ),
330 _ => Err(Reject::Protocol),
331 }
332}