Skip to main content

heddle_thread_api/transport/
iroh.rs

1// SPDX-License-Identifier: Apache-2.0
2//! One reliable Iroh bidirectional stream per RPC, with bounded framing and
3//! cancellation-safe readers. The caller owns connection and credential setup.
4use 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
120/// Adapt an already accepted RPC stream after the caller parsed its bounded
121/// request prelude. This performs framing only; the server must independently
122/// authorize the exact opening before admission or disclosure.
123pub 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            // Count the actual boundary read even if this chunk exceeds budget.
155            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            // Keep a started frame's deadline on Reader, just like its bytes.
194            // Cancellation and a later next() cannot restart that deadline.
195            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                // Check before allocating or reading any body bytes.
211                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            // Partial framing lives on Reader, so selecting another future while
228            // next() waits cannot discard a consumed header or body prefix.
229            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                // A live stream may have no new records for hours. Iroh owns
239                // connection liveness; callers own cancellation/overall waits.
240                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    /// Send a bounded typed failure and close this accepted response stream.
303    /// A rejected oversized failure leaves the stream available for a smaller
304    /// failure; a successful terminal failure permits no subsequent writes.
305    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;