heddle_thread_api/transport/
iroh.rs1use std::time::Duration;
5
6use ::iroh::endpoint::{Connection, RecvStream, SendStream};
7use api::{
8 framing,
9 heddle::api::common::CallFailure,
10 v2::{
11 MethodDescriptor,
12 client::{MessageReader, MessageWriter, RpcTransport},
13 },
14};
15use prost::Message;
16
17use super::{Authorize, Error};
18
19fn io(error: impl std::fmt::Display) -> Error {
20 Error::Io(error.to_string())
21}
22
23pub struct IrohTransport<A> {
24 connection: Connection,
25 authorize: A,
26 frame_limit: usize,
27 progress_timeout: Duration,
28}
29
30impl<A: Authorize> IrohTransport<A> {
31 pub fn new(
32 connection: Connection,
33 authorize: A,
34 frame_limit: usize,
35 progress_timeout: Duration,
36 ) -> Result<Self, Error> {
37 if frame_limit == 0 || frame_limit > framing::MAX_CONTROL_BODY || progress_timeout.is_zero()
38 {
39 return Err(Error::Protocol("invalid local transport limits"));
40 }
41 Ok(Self {
42 connection,
43 authorize,
44 frame_limit,
45 progress_timeout,
46 })
47 }
48
49 async fn open(
50 &self,
51 method: &'static MethodDescriptor,
52 body: &[u8],
53 exchange: bool,
54 ) -> Result<(Writer, Reader), Error> {
55 if body.len() > self.frame_limit {
56 return Err(Error::Protocol("request exceeds frame budget"));
57 }
58 let mut context = self.authorize.context(method, body).await?;
59 context.client_operation_id = method.client_operation_id(body)?.unwrap_or_default().into();
60 let (send, recv) = tokio::time::timeout(self.progress_timeout, self.connection.open_bi())
61 .await
62 .map_err(|_| Error::Timeout)?
63 .map_err(io)?;
64 heddle_perf_contract::record_network_stream_opened();
65 let mut writer = Writer::new(send, self.frame_limit, self.progress_timeout);
66 let reader = Reader::for_method(recv, self.frame_limit, self.progress_timeout, method);
67 let frame = if exchange {
68 framing::encode_request_prelude(method.path, &context)?
69 } else {
70 framing::encode_request_frame(method.path, &context, body)?
71 };
72 writer.write(&frame).await?;
73 if exchange {
74 writer.send(body.to_vec()).await?;
75 } else {
76 writer.finish().await?;
77 }
78 Ok((writer, reader))
79 }
80}
81
82impl<A: Authorize> RpcTransport for IrohTransport<A> {
83 type Error = Error;
84 type Reader = Reader;
85 type Writer = Writer;
86
87 async fn unary(
88 &self,
89 method: &'static MethodDescriptor,
90 request: Vec<u8>,
91 ) -> Result<Vec<u8>, Error> {
92 let (_, mut reader) = self.open(method, &request, false).await?;
93 let frame = tokio::time::timeout(reader.timeout, reader.read_unary())
94 .await
95 .map_err(|_| Error::Timeout)??;
96 reader.done = true;
97 match framing::decode_response_frame(&frame)? {
98 framing::ResponseFrame::Success(bytes) => Ok(bytes.to_vec()),
99 framing::ResponseFrame::Failure(failure) => Err(Error::Remote(failure.into())),
100 }
101 }
102
103 async fn observe(
104 &self,
105 method: &'static MethodDescriptor,
106 request: Vec<u8>,
107 ) -> Result<Reader, Error> {
108 Ok(self.open(method, &request, false).await?.1)
109 }
110
111 async fn exchange(
112 &self,
113 method: &'static MethodDescriptor,
114 opening: Vec<u8>,
115 ) -> Result<(Writer, Reader), Error> {
116 self.open(method, &opening, true).await
117 }
118}
119
120pub fn accepted_stream(
124 send: SendStream,
125 recv: RecvStream,
126 frame_limit: usize,
127 progress_timeout: Duration,
128 method: &MethodDescriptor,
129) -> Result<(Writer, Reader), Error> {
130 if frame_limit == 0 || frame_limit > framing::MAX_CONTROL_BODY || progress_timeout.is_zero() {
131 return Err(Error::Protocol("invalid accepted transport limits"));
132 }
133 Ok((
134 Writer::new(send, frame_limit, progress_timeout),
135 Reader::for_method(recv, frame_limit, progress_timeout, method),
136 ))
137}
138
139pub struct Reader {
140 recv: RecvStream,
141 frame_limit: usize,
142 timeout: Duration,
143 done: bool,
144 buffer: Vec<u8>,
145 opening: bool,
146 live: bool,
147 frame_deadline: Option<tokio::time::Instant>,
148}
149
150impl Reader {
151 async fn read_unary(&mut self) -> Result<Vec<u8>, Error> {
152 let mut body = Vec::new();
153 while let Some(chunk) = self.recv.read_chunk(8192).await.map_err(io)? {
154 heddle_perf_contract::record_network_bytes_received(chunk.len());
156 if body.len().saturating_add(chunk.len()) > self.frame_limit + 1 {
157 return Err(Error::Protocol("response exceeds frame budget"));
158 }
159 body.extend_from_slice(&chunk);
160 }
161 Ok(body)
162 }
163
164 pub(crate) fn new(recv: RecvStream, frame_limit: usize, timeout: Duration) -> Self {
165 Self {
166 recv,
167 frame_limit,
168 timeout,
169 done: false,
170 buffer: Vec::with_capacity(5),
171 opening: true,
172 live: false,
173 frame_deadline: None,
174 }
175 }
176
177 pub(crate) fn for_method(
178 recv: RecvStream,
179 frame_limit: usize,
180 timeout: Duration,
181 method: &MethodDescriptor,
182 ) -> Self {
183 let mut reader = Self::new(recv, frame_limit, timeout);
184 reader.live = method.live_stream;
185 reader
186 }
187
188 async fn read_frame(&mut self) -> Result<Option<Vec<u8>>, Error> {
189 if (!self.live || self.opening) && self.frame_deadline.is_none() {
190 self.frame_deadline = Some(tokio::time::Instant::now() + self.timeout);
191 }
192 loop {
193 if self
196 .frame_deadline
197 .is_some_and(|deadline| tokio::time::Instant::now() >= deadline)
198 {
199 return Err(Error::Timeout);
200 }
201 let needed = if self.buffer.len() < 5 {
202 5
203 } else {
204 let header = &self.buffer;
205 if header[0] > 1 {
206 return Err(Error::Protocol("unexpected stream frame kind"));
207 }
208 let length =
209 u32::from_be_bytes([header[1], header[2], header[3], header[4]]) as usize;
210 if length > self.frame_limit {
212 return Err(Error::Protocol("response exceeds frame budget"));
213 }
214 length + 5
215 };
216 if self.buffer.len() == needed {
217 let mut frame = std::mem::take(&mut self.buffer);
218 let kind = frame[0];
219 frame.drain(..5);
220 if kind == 1 {
221 return Err(Error::Remote(CallFailure::decode(frame.as_slice())?.into()));
222 }
223 self.opening = false;
224 self.frame_deadline = None;
225 return Ok(Some(frame));
226 }
227 let mut chunk = [0; 8192];
230 let size = (needed - self.buffer.len()).min(chunk.len());
231 let read = match self.frame_deadline {
232 Some(deadline) => {
233 tokio::time::timeout_at(deadline, self.recv.read(&mut chunk[..size]))
234 .await
235 .map_err(|_| Error::Timeout)?
236 .map_err(io)?
237 }
238 None => self.recv.read(&mut chunk[..size]).await.map_err(io)?,
241 };
242 match read {
243 Some(n) => {
244 heddle_perf_contract::record_network_bytes_received(n);
245 if self.frame_deadline.is_none() {
246 self.frame_deadline = Some(tokio::time::Instant::now() + self.timeout);
247 }
248 self.buffer.extend_from_slice(&chunk[..n]);
249 }
250 None if self.buffer.is_empty() => {
251 self.done = true;
252 return Ok(None);
253 }
254 None => return Err(Error::Protocol("FIN within a stream frame")),
255 }
256 }
257 }
258}
259
260impl MessageReader for Reader {
261 type Error = Error;
262 async fn next(&mut self) -> Result<Option<Vec<u8>>, Error> {
263 if self.done {
264 return Ok(None);
265 }
266 let result = self.read_frame().await;
267 if result.is_err() {
268 self.cancel();
269 }
270 result
271 }
272 fn cancel(&mut self) {
273 if !self.done {
274 let _ = self.recv.stop(0u32.into());
275 self.done = true;
276 self.buffer = Vec::new();
277 }
278 }
279}
280impl Drop for Reader {
281 fn drop(&mut self) {
282 self.cancel();
283 }
284}
285
286pub struct Writer {
287 send: Option<SendStream>,
288 frame_limit: usize,
289 timeout: Duration,
290 poisoned: bool,
291}
292impl Writer {
293 pub(crate) fn new(send: SendStream, frame_limit: usize, timeout: Duration) -> Self {
294 Self {
295 send: Some(send),
296 frame_limit,
297 timeout,
298 poisoned: false,
299 }
300 }
301
302 pub async fn fail(&mut self, failure: &CallFailure) -> Result<(), Error> {
306 if failure.encoded_len() > self.frame_limit {
307 return Err(Error::Protocol("failure exceeds frame budget"));
308 }
309 self.write(&framing::encode_stream_failure(failure)?)
310 .await?;
311 self.finish().await
312 }
313
314 async fn write(&mut self, bytes: &[u8]) -> Result<(), Error> {
315 let send = self
316 .send
317 .as_mut()
318 .ok_or(Error::Protocol("request stream is closed"))?;
319 if self.poisoned {
320 return Err(Error::Protocol(
321 "previous write interrupted; reopen the RPC",
322 ));
323 }
324 self.poisoned = true;
325 tokio::time::timeout(self.timeout, async {
326 let mut remaining = bytes;
327 while !remaining.is_empty() {
328 let written = send.write(remaining).await.map_err(io)?;
329 heddle_perf_contract::record_network_bytes_sent(written);
330 remaining = &remaining[written..];
331 }
332 Ok::<_, Error>(())
333 })
334 .await
335 .map_err(|_| Error::Timeout)??;
336 self.poisoned = false;
337 Ok(())
338 }
339}
340impl MessageWriter for Writer {
341 type Error = Error;
342 async fn send(&mut self, message: Vec<u8>) -> Result<(), Error> {
343 if message.len() > self.frame_limit {
344 return Err(Error::Protocol("request exceeds frame budget"));
345 }
346 self.write(&framing::encode_stream_message(&message)?).await
347 }
348 async fn finish(&mut self) -> Result<(), Error> {
349 if self.poisoned {
350 self.abort();
351 return Err(Error::Protocol("cannot finish an interrupted request"));
352 }
353 if let Some(mut send) = self.send.take() {
354 send.finish().map_err(io)?;
355 }
356 Ok(())
357 }
358 fn abort(&mut self) {
359 if let Some(mut send) = self.send.take() {
360 let _ = send.reset(0u32.into());
361 }
362 }
363}
364impl Drop for Writer {
365 fn drop(&mut self) {
366 self.abort();
367 }
368}
369
370#[cfg(test)]
371#[path = "../transport_tests.rs"]
372mod tests;