1#[cfg(feature = "quic")]
8use std::time::Duration;
9use std::{
10 collections::VecDeque,
11 future::Future,
12 net::SocketAddr,
13 sync::{Arc, Mutex, OnceLock, mpsc},
14};
15
16use bytes::{Buf, BytesMut};
17use datum::{
18 NotUsed, Sink, Source, SourceRef, StreamCompletion, StreamError, StreamRefFrame, StreamRefId,
19 StreamRefMessage, StreamRefOutbound, StreamRefPayload, StreamRefPayloadBatch,
20 StreamRefProtoConsumer, StreamRefProtoEndpoint, StreamRefProtoProducer, StreamRefSettings,
21 StreamResult,
22 actor::stream_ref_proto::{StreamRefOutboundPoll, StreamRefProtoEndpointWake},
23};
24#[cfg(feature = "quic")]
25use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite};
26use tokio::{
27 io::AsyncWriteExt,
28 net::{TcpListener, TcpStream, ToSocketAddrs},
29 runtime::{Handle, Runtime},
30 sync::mpsc as tokio_mpsc,
31 task::JoinHandle,
32};
33
34#[cfg(feature = "quic")]
35use crate::QuicBidirectionalStream;
36
37const FRAME_LEN_BYTES: usize = 4;
38const MAX_STREAM_REF_FRAME_BYTES: usize = 16 * 1024 * 1024;
39const STREAM_REF_TCP_CHUNK_SIZE: usize = 8192;
40#[cfg(feature = "quic")]
41const STREAM_REF_QUIC_READ_BUFFER_BYTES: usize = 2048;
42const STREAM_REF_OUTBOUND_BATCH_FRAMES: usize = 64;
43#[cfg(feature = "quic")]
44const STREAM_REF_OUTBOUND_RECHECK_INTERVAL: Duration = Duration::from_millis(5);
45
46const COMPACT_FRAME_FLAG: u32 = 0x8000_0000;
50const COMPACT_FRAME_LEN_MASK: u32 = 0x7fff_ffff;
51const COMPACT_FRAME_VERSION: u8 = 1;
52const COMPACT_SEQUENCED_ON_NEXT_BATCH: u8 = 1;
53const COMPACT_BATCH_HEADER_BYTES: usize = 1 + 1 + 16 + 8 + 2;
54const COMPACT_BATCH_ELEMENT_LEN_BYTES: usize = 4;
55
56#[derive(Clone, Copy)]
57struct CarrierReadMode {
58 chunk_size: usize,
59 emit_available: bool,
60 fail_on_eof: bool,
61}
62
63impl CarrierReadMode {
64 fn new(chunk_size: usize, emit_available: bool, fail_on_eof: bool) -> Self {
65 assert!(chunk_size > 0, "chunk size must be greater than zero");
66 Self {
67 chunk_size,
68 emit_available,
69 fail_on_eof,
70 }
71 }
72}
73
74#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
77pub struct StreamRefProtocolMessageCounts {
78 pub cumulative_demand: u64,
79 pub sequenced_on_next: u64,
80 pub ack: u64,
81}
82
83#[derive(Clone, Default)]
89pub struct StreamRefProtocolDiagnostics {
90 counts: Arc<Mutex<StreamRefProtocolMessageCounts>>,
91}
92
93impl StreamRefProtocolDiagnostics {
94 #[must_use]
95 pub fn new() -> Self {
96 Self::default()
97 }
98
99 #[must_use]
100 pub fn snapshot(&self) -> StreamRefProtocolMessageCounts {
101 *self
102 .counts
103 .lock()
104 .unwrap_or_else(|poison| poison.into_inner())
105 }
106
107 fn record_counts(&self, delta: StreamRefProtocolMessageCounts) {
108 if delta == StreamRefProtocolMessageCounts::default() {
109 return;
110 }
111 let mut counts = self
112 .counts
113 .lock()
114 .unwrap_or_else(|poison| poison.into_inner());
115 counts.cumulative_demand = counts
116 .cumulative_demand
117 .saturating_add(delta.cumulative_demand);
118 counts.sequenced_on_next = counts
119 .sequenced_on_next
120 .saturating_add(delta.sequenced_on_next);
121 counts.ack = counts.ack.saturating_add(delta.ack);
122 }
123}
124
125fn outbound_counts(outbound: &StreamRefOutbound) -> StreamRefProtocolMessageCounts {
126 let mut counts = StreamRefProtocolMessageCounts::default();
127 match outbound {
128 StreamRefOutbound::Frame(frame) => match &frame.message {
129 StreamRefMessage::CumulativeDemand { .. } => {
130 counts.cumulative_demand = 1;
131 }
132 StreamRefMessage::SequencedOnNext { .. } => {
133 counts.sequenced_on_next = 1;
134 }
135 StreamRefMessage::Ack => {
136 counts.ack = 1;
137 }
138 StreamRefMessage::OnSubscribeHandshake
139 | StreamRefMessage::RemoteStreamCompleted { .. }
140 | StreamRefMessage::RemoteStreamFailure { .. } => {}
141 },
142 StreamRefOutbound::SequencedBatch(batch) => {
143 counts.sequenced_on_next = batch.count() as u64;
144 }
145 }
146 counts
147}
148
149#[derive(Clone, Copy)]
150struct PendingDiagnostic {
151 remaining: usize,
152 counts: StreamRefProtocolMessageCounts,
153}
154
155#[must_use = "wait for the QUIC StreamRefs carrier to observe completion or failure"]
157#[cfg(feature = "quic")]
158pub struct StreamRefQuicHandle {
159 completion: EndpointTaskCompletion,
160}
161
162#[cfg(feature = "quic")]
163impl StreamRefQuicHandle {
164 pub fn wait(self) -> StreamResult<NotUsed> {
165 self.completion.wait()
166 }
167
168 #[must_use]
169 pub fn try_wait(&self) -> Option<StreamResult<NotUsed>> {
170 self.completion.try_wait()
171 }
172
173 #[must_use]
174 pub fn is_finished(&self) -> bool {
175 self.completion.is_finished()
176 }
177}
178
179#[derive(Debug, Clone, Copy, PartialEq, Eq)]
181#[cfg(feature = "tcp")]
182pub struct StreamRefTcpBinding {
183 local_addr: SocketAddr,
184}
185
186#[cfg(feature = "tcp")]
187impl StreamRefTcpBinding {
188 #[must_use]
189 pub fn local_addr(&self) -> SocketAddr {
190 self.local_addr
191 }
192}
193
194#[must_use = "wait for the TCP StreamRefs carrier to observe completion or failure"]
196#[cfg(feature = "tcp")]
197pub struct StreamRefTcpHandle {
198 completion: EndpointTaskCompletion,
199}
200
201#[cfg(feature = "tcp")]
202impl StreamRefTcpHandle {
203 pub fn wait(self) -> StreamResult<NotUsed> {
204 self.completion.wait()
205 }
206
207 #[must_use]
208 pub fn try_wait(&self) -> Option<StreamResult<NotUsed>> {
209 self.completion.try_wait()
210 }
211
212 #[must_use]
213 pub fn is_finished(&self) -> bool {
214 self.completion.is_finished()
215 }
216}
217
218struct EndpointTaskCompletion {
219 receiver: mpsc::Receiver<StreamResult<NotUsed>>,
220 task: Option<JoinHandle<()>>,
221}
222
223impl EndpointTaskCompletion {
224 fn wait(mut self) -> StreamResult<NotUsed> {
225 let result = self
226 .receiver
227 .recv()
228 .unwrap_or(Err(StreamError::AbruptTermination));
229 self.task.take();
230 result
231 }
232
233 fn try_wait(&self) -> Option<StreamResult<NotUsed>> {
234 self.receiver.try_recv().ok()
235 }
236
237 fn is_finished(&self) -> bool {
238 self.task.as_ref().is_some_and(JoinHandle::is_finished)
239 }
240}
241
242impl Drop for EndpointTaskCompletion {
243 fn drop(&mut self) {
244 if let Some(task) = self.task.take() {
245 task.abort();
246 }
247 }
248}
249
250#[cfg(feature = "quic")]
252pub fn serve_source_ref_over_quic<T>(
253 stream: QuicBidirectionalStream,
254 source_ref: SourceRef<T>,
255 stream_ref_id: StreamRefId,
256 settings: StreamRefSettings,
257) -> StreamResult<StreamRefQuicHandle>
258where
259 T: StreamRefPayload,
260{
261 let producer = StreamRefProtoProducer::from_source_ref(source_ref, stream_ref_id, settings)?;
262 Ok(drive_stream_ref_endpoint_over_quic(stream, producer, None))
263}
264
265#[cfg(feature = "quic")]
267pub fn serve_source_over_quic<T, Mat>(
268 stream: QuicBidirectionalStream,
269 source: Source<T, Mat>,
270 stream_ref_id: StreamRefId,
271 settings: StreamRefSettings,
272) -> StreamResult<StreamRefQuicHandle>
273where
274 T: StreamRefPayload,
275 Mat: Send + 'static,
276{
277 let producer = StreamRefProtoProducer::from_source(source, stream_ref_id, settings)?;
278 Ok(drive_stream_ref_endpoint_over_quic(stream, producer, None))
279}
280
281#[cfg(feature = "quic")]
283pub fn source_ref_over_quic<T>(
284 stream: QuicBidirectionalStream,
285 stream_ref_id: StreamRefId,
286 settings: StreamRefSettings,
287) -> (Source<T, NotUsed>, StreamRefQuicHandle)
288where
289 T: StreamRefPayload,
290{
291 let consumer = StreamRefProtoConsumer::new(stream_ref_id, settings);
292 let source = consumer.source();
293 let handle = drive_stream_ref_endpoint_over_quic(stream, consumer, None);
294 (source, handle)
295}
296
297#[cfg(feature = "quic")]
305pub fn serve_sink_ref_over_quic<T>(
306 stream: QuicBidirectionalStream,
307 stream_ref_id: StreamRefId,
308 settings: StreamRefSettings,
309) -> (Source<T, NotUsed>, StreamRefQuicHandle)
310where
311 T: StreamRefPayload,
312{
313 let consumer = StreamRefProtoConsumer::new(stream_ref_id, settings);
314 let source = consumer.source();
315 let handle = drive_stream_ref_endpoint_over_quic(stream, consumer, None);
316 (source, handle)
317}
318
319#[cfg(feature = "quic")]
322pub fn sink_ref_over_quic<T>(
323 stream: QuicBidirectionalStream,
324 stream_ref_id: StreamRefId,
325 settings: StreamRefSettings,
326) -> (Sink<T, StreamCompletion<NotUsed>>, StreamRefQuicHandle)
327where
328 T: StreamRefPayload,
329{
330 let producer = StreamRefProtoProducer::new_lazy(stream_ref_id, settings);
331 let sink = producer.sink();
332 let handle = drive_stream_ref_endpoint_over_quic(stream, producer, None);
333 (sink, handle)
334}
335
336#[cfg(feature = "tcp")]
343pub fn serve_source_ref_over_tcp<T, A>(
344 addr: A,
345 source_ref: SourceRef<T>,
346 stream_ref_id: StreamRefId,
347 settings: StreamRefSettings,
348) -> StreamResult<(StreamRefTcpBinding, StreamRefTcpHandle)>
349where
350 T: StreamRefPayload,
351 A: ToSocketAddrs + Send + 'static,
352{
353 serve_source_ref_over_tcp_with_diagnostics(addr, source_ref, stream_ref_id, settings, None)
354}
355
356#[cfg(feature = "tcp")]
357pub fn serve_source_ref_over_tcp_with_diagnostics<T, A>(
358 addr: A,
359 source_ref: SourceRef<T>,
360 stream_ref_id: StreamRefId,
361 settings: StreamRefSettings,
362 diagnostics: Option<StreamRefProtocolDiagnostics>,
363) -> StreamResult<(StreamRefTcpBinding, StreamRefTcpHandle)>
364where
365 T: StreamRefPayload,
366 A: ToSocketAddrs + Send + 'static,
367{
368 let producer = StreamRefProtoProducer::from_source_ref(source_ref, stream_ref_id, settings)?;
369 let (listener, binding, handle) = bind_tcp_listener(addr)?;
370 Ok((
371 binding,
372 drive_stream_ref_endpoint_over_tcp_listener(listener, handle, producer, diagnostics),
373 ))
374}
375
376#[cfg(feature = "tcp")]
382pub fn serve_source_ref_over_tcp_stream<T>(
383 stream: TcpStream,
384 source_ref: SourceRef<T>,
385 stream_ref_id: StreamRefId,
386 settings: StreamRefSettings,
387) -> StreamResult<StreamRefTcpHandle>
388where
389 T: StreamRefPayload,
390{
391 serve_source_ref_over_tcp_stream_with_diagnostics(
392 stream,
393 source_ref,
394 stream_ref_id,
395 settings,
396 None,
397 )
398}
399
400#[cfg(feature = "tcp")]
401pub fn serve_source_ref_over_tcp_stream_with_diagnostics<T>(
402 stream: TcpStream,
403 source_ref: SourceRef<T>,
404 stream_ref_id: StreamRefId,
405 settings: StreamRefSettings,
406 diagnostics: Option<StreamRefProtocolDiagnostics>,
407) -> StreamResult<StreamRefTcpHandle>
408where
409 T: StreamRefPayload,
410{
411 let producer = StreamRefProtoProducer::from_source_ref(source_ref, stream_ref_id, settings)?;
412 let handle = current_tokio_handle()?;
413 Ok(drive_stream_ref_endpoint_over_tcp_stream(
414 stream,
415 handle,
416 producer,
417 diagnostics,
418 ))
419}
420
421#[cfg(feature = "tcp")]
426pub fn source_ref_over_tcp<T, A>(
427 addr: A,
428 stream_ref_id: StreamRefId,
429 settings: StreamRefSettings,
430) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
431where
432 T: StreamRefPayload,
433 A: ToSocketAddrs + Send + 'static,
434{
435 source_ref_over_tcp_with_diagnostics(addr, stream_ref_id, settings, None)
436}
437
438#[cfg(feature = "tcp")]
439pub fn source_ref_over_tcp_with_diagnostics<T, A>(
440 addr: A,
441 stream_ref_id: StreamRefId,
442 settings: StreamRefSettings,
443 diagnostics: Option<StreamRefProtocolDiagnostics>,
444) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
445where
446 T: StreamRefPayload,
447 A: ToSocketAddrs + Send + 'static,
448{
449 let consumer = StreamRefProtoConsumer::new(stream_ref_id, settings);
450 let source = consumer.source();
451 let (stream, handle) = connect_tcp_stream(addr)?;
452 let handle = drive_stream_ref_endpoint_over_tcp_stream(stream, handle, consumer, diagnostics);
453 Ok((source, handle))
454}
455
456#[cfg(feature = "tcp")]
459pub fn source_ref_over_tcp_stream<T>(
460 stream: TcpStream,
461 stream_ref_id: StreamRefId,
462 settings: StreamRefSettings,
463) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
464where
465 T: StreamRefPayload,
466{
467 source_ref_over_tcp_stream_with_diagnostics(stream, stream_ref_id, settings, None)
468}
469
470#[cfg(feature = "tcp")]
471pub fn source_ref_over_tcp_stream_with_diagnostics<T>(
472 stream: TcpStream,
473 stream_ref_id: StreamRefId,
474 settings: StreamRefSettings,
475 diagnostics: Option<StreamRefProtocolDiagnostics>,
476) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
477where
478 T: StreamRefPayload,
479{
480 let consumer = StreamRefProtoConsumer::new(stream_ref_id, settings);
481 let source = consumer.source();
482 let handle = current_tokio_handle()?;
483 let handle = drive_stream_ref_endpoint_over_tcp_stream(stream, handle, consumer, diagnostics);
484 Ok((source, handle))
485}
486
487#[cfg(feature = "tcp")]
493pub fn serve_sink_ref_over_tcp<T, A>(
494 addr: A,
495 stream_ref_id: StreamRefId,
496 settings: StreamRefSettings,
497) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
498where
499 T: StreamRefPayload,
500 A: ToSocketAddrs + Send + 'static,
501{
502 serve_sink_ref_over_tcp_with_diagnostics(addr, stream_ref_id, settings, None)
503}
504
505#[cfg(feature = "tcp")]
506pub fn serve_sink_ref_over_tcp_with_diagnostics<T, A>(
507 addr: A,
508 stream_ref_id: StreamRefId,
509 settings: StreamRefSettings,
510 diagnostics: Option<StreamRefProtocolDiagnostics>,
511) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
512where
513 T: StreamRefPayload,
514 A: ToSocketAddrs + Send + 'static,
515{
516 let consumer = StreamRefProtoConsumer::new(stream_ref_id, settings);
517 let source = consumer.source();
518 let (stream, handle) = connect_tcp_stream(addr)?;
519 let handle = drive_stream_ref_endpoint_over_tcp_stream(stream, handle, consumer, diagnostics);
520 Ok((source, handle))
521}
522
523#[cfg(feature = "tcp")]
526pub fn serve_sink_ref_over_tcp_stream<T>(
527 stream: TcpStream,
528 stream_ref_id: StreamRefId,
529 settings: StreamRefSettings,
530) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
531where
532 T: StreamRefPayload,
533{
534 serve_sink_ref_over_tcp_stream_with_diagnostics(stream, stream_ref_id, settings, None)
535}
536
537#[cfg(feature = "tcp")]
538pub fn serve_sink_ref_over_tcp_stream_with_diagnostics<T>(
539 stream: TcpStream,
540 stream_ref_id: StreamRefId,
541 settings: StreamRefSettings,
542 diagnostics: Option<StreamRefProtocolDiagnostics>,
543) -> StreamResult<(Source<T, NotUsed>, StreamRefTcpHandle)>
544where
545 T: StreamRefPayload,
546{
547 let consumer = StreamRefProtoConsumer::new(stream_ref_id, settings);
548 let source = consumer.source();
549 let handle = current_tokio_handle()?;
550 let handle = drive_stream_ref_endpoint_over_tcp_stream(stream, handle, consumer, diagnostics);
551 Ok((source, handle))
552}
553
554#[cfg(feature = "tcp")]
562pub fn sink_ref_over_tcp<T, A>(
563 addr: A,
564 stream_ref_id: StreamRefId,
565 settings: StreamRefSettings,
566) -> StreamResult<(
567 Sink<T, StreamCompletion<NotUsed>>,
568 StreamRefTcpBinding,
569 StreamRefTcpHandle,
570)>
571where
572 T: StreamRefPayload,
573 A: ToSocketAddrs + Send + 'static,
574{
575 sink_ref_over_tcp_with_diagnostics(addr, stream_ref_id, settings, None)
576}
577
578#[cfg(feature = "tcp")]
579pub fn sink_ref_over_tcp_with_diagnostics<T, A>(
580 addr: A,
581 stream_ref_id: StreamRefId,
582 settings: StreamRefSettings,
583 diagnostics: Option<StreamRefProtocolDiagnostics>,
584) -> StreamResult<(
585 Sink<T, StreamCompletion<NotUsed>>,
586 StreamRefTcpBinding,
587 StreamRefTcpHandle,
588)>
589where
590 T: StreamRefPayload,
591 A: ToSocketAddrs + Send + 'static,
592{
593 let producer = StreamRefProtoProducer::new_lazy(stream_ref_id, settings);
594 let sink = producer.sink();
595 let (listener, binding, handle) = bind_tcp_listener(addr)?;
596 let handle =
597 drive_stream_ref_endpoint_over_tcp_listener(listener, handle, producer, diagnostics);
598 Ok((sink, binding, handle))
599}
600
601#[cfg(feature = "tcp")]
604pub fn sink_ref_over_tcp_stream<T>(
605 stream: TcpStream,
606 stream_ref_id: StreamRefId,
607 settings: StreamRefSettings,
608) -> StreamResult<(Sink<T, StreamCompletion<NotUsed>>, StreamRefTcpHandle)>
609where
610 T: StreamRefPayload,
611{
612 sink_ref_over_tcp_stream_with_diagnostics(stream, stream_ref_id, settings, None)
613}
614
615#[cfg(feature = "tcp")]
616pub fn sink_ref_over_tcp_stream_with_diagnostics<T>(
617 stream: TcpStream,
618 stream_ref_id: StreamRefId,
619 settings: StreamRefSettings,
620 diagnostics: Option<StreamRefProtocolDiagnostics>,
621) -> StreamResult<(Sink<T, StreamCompletion<NotUsed>>, StreamRefTcpHandle)>
622where
623 T: StreamRefPayload,
624{
625 let producer = StreamRefProtoProducer::new_lazy(stream_ref_id, settings);
626 let sink = producer.sink();
627 let handle = current_tokio_handle()?;
628 let handle = drive_stream_ref_endpoint_over_tcp_stream(stream, handle, producer, diagnostics);
629 Ok((sink, handle))
630}
631
632#[cfg(feature = "quic")]
633fn drive_stream_ref_endpoint_over_quic<E>(
634 stream: QuicBidirectionalStream,
635 endpoint: E,
636 diagnostics: Option<StreamRefProtocolDiagnostics>,
637) -> StreamRefQuicHandle
638where
639 E: StreamRefProtoEndpointWake,
640{
641 let (reader, writer, handle, chunk_size, emit_available) = stream.into_stream_ref_parts();
642 let read_mode = CarrierReadMode::new(chunk_size, emit_available, false);
643 StreamRefQuicHandle {
644 completion: spawn_endpoint_task(&handle, async move {
645 run_stream_ref_endpoint_quic_task(reader, writer, endpoint, read_mode, diagnostics)
646 .await
647 }),
648 }
649}
650
651#[cfg(feature = "tcp")]
652fn drive_stream_ref_endpoint_over_tcp_listener<E>(
653 listener: TcpListener,
654 handle: Handle,
655 endpoint: E,
656 diagnostics: Option<StreamRefProtocolDiagnostics>,
657) -> StreamRefTcpHandle
658where
659 E: StreamRefProtoEndpointWake,
660{
661 StreamRefTcpHandle {
662 completion: spawn_endpoint_task(&handle, async move {
663 let (stream, _) = listener.accept().await.map_err(io_error)?;
664 run_stream_ref_endpoint_tcp_task(stream, endpoint, diagnostics).await
665 }),
666 }
667}
668
669#[cfg(feature = "tcp")]
670fn drive_stream_ref_endpoint_over_tcp_stream<E>(
671 stream: TcpStream,
672 handle: Handle,
673 endpoint: E,
674 diagnostics: Option<StreamRefProtocolDiagnostics>,
675) -> StreamRefTcpHandle
676where
677 E: StreamRefProtoEndpointWake,
678{
679 StreamRefTcpHandle {
680 completion: spawn_endpoint_task(&handle, async move {
681 run_stream_ref_endpoint_tcp_task(stream, endpoint, diagnostics).await
682 }),
683 }
684}
685
686#[cfg(feature = "quic")]
687async fn run_stream_ref_endpoint_quic_task<R, W, E>(
688 reader: R,
689 writer: W,
690 endpoint: E,
691 read_mode: CarrierReadMode,
692 diagnostics: Option<StreamRefProtocolDiagnostics>,
693) -> StreamResult<NotUsed>
694where
695 R: AsyncRead + Unpin + Send + 'static,
696 W: AsyncWrite + Unpin + Send + 'static,
697 E: StreamRefProtoEndpointWake,
698{
699 let (wake_sender, wake_receiver) = tokio_mpsc::channel(1);
700 endpoint.install_outbound_wake(wake_sender.clone());
701 let _ = wake_sender.try_send(());
702
703 let result = QuicEndpointTask {
704 reader,
705 writer,
706 endpoint: endpoint.clone(),
707 diagnostics,
708 read_mode,
709 decoder: FrameDecoder::default(),
710 read_buffer: vec![
711 0_u8;
712 read_mode
713 .chunk_size
714 .clamp(1, STREAM_REF_QUIC_READ_BUFFER_BYTES)
715 ],
716 pending_tail: Vec::new(),
717 write_buffer: BytesMut::new(),
718 encode_buffer: Vec::new(),
719 pending_diagnostics: VecDeque::new(),
720 read_closed: false,
721 inbound_seen: false,
722 outbound_written: false,
723 recheck_outbound: false,
724 outbound_closed: false,
725 write_shutdown: false,
726 wake_receiver,
727 }
728 .run()
729 .await;
730
731 endpoint.clear_outbound_wake();
732 if let Err(error) = &result {
733 endpoint.fail_connection(error.clone());
734 }
735 result
736}
737
738#[cfg(feature = "quic")]
739struct QuicEndpointTask<R, W, E>
740where
741 E: StreamRefProtoEndpointWake,
742{
743 reader: R,
744 writer: W,
745 endpoint: E,
746 diagnostics: Option<StreamRefProtocolDiagnostics>,
747 read_mode: CarrierReadMode,
748 decoder: FrameDecoder,
749 read_buffer: Vec<u8>,
750 pending_tail: Vec<u8>,
751 write_buffer: BytesMut,
752 encode_buffer: Vec<u8>,
753 pending_diagnostics: VecDeque<PendingDiagnostic>,
754 read_closed: bool,
755 inbound_seen: bool,
756 outbound_written: bool,
757 recheck_outbound: bool,
758 outbound_closed: bool,
759 write_shutdown: bool,
760 wake_receiver: tokio_mpsc::Receiver<()>,
761}
762
763#[cfg(feature = "quic")]
764impl<R, W, E> QuicEndpointTask<R, W, E>
765where
766 R: AsyncRead + Unpin,
767 W: AsyncWrite + Unpin,
768 E: StreamRefProtoEndpointWake,
769{
770 async fn run(mut self) -> StreamResult<NotUsed> {
771 loop {
772 if self.read_closed && self.write_shutdown {
773 return Ok(NotUsed);
774 }
775
776 self.drain_outbound()?;
777 if self.read_closed
778 && self.outbound_written
779 && !self.outbound_closed
780 && self.write_buffer.is_empty()
781 && self.recheck_outbound
782 {
783 return Err(StreamError::Failed(
784 "StreamRefs QUIC peer closed before remote terminal".to_owned(),
785 ));
786 }
787 if self.outbound_closed && self.write_buffer.is_empty() && !self.write_shutdown {
788 self.writer.shutdown().await.map_err(io_error)?;
789 self.write_shutdown = true;
790 continue;
791 }
792 tokio::select! {
793 biased;
794 wake = self.wake_receiver.recv(), if !self.outbound_closed => {
795 if wake.is_none() {
796 self.drain_outbound()?;
797 }
798 }
799 _ = tokio::time::sleep(STREAM_REF_OUTBOUND_RECHECK_INTERVAL), if self.recheck_outbound && !self.outbound_closed && self.write_buffer.is_empty() => {
800 self.drain_outbound()?;
801 }
802 written = self.writer.write(&self.write_buffer), if !self.write_buffer.is_empty() => {
803 self.handle_written(written)?;
804 }
805 read = self.reader.read(&mut self.read_buffer), if !self.read_closed => {
806 match read {
807 Ok(0) => self.handle_eof()?,
808 Ok(read) => {
809 self.feed_read_buffer(read)?;
810 self.drain_outbound()?;
811 }
812 Err(error) => {
813 let error = io_error(error);
814 if self.write_shutdown && is_quic_teardown_loss(&error) {
815 return Ok(NotUsed);
816 }
817 return Err(error);
818 }
819 }
820 }
821 }
822 }
823 }
824
825 fn drain_outbound(&mut self) -> StreamResult<()> {
826 self.recheck_outbound = false;
827 while !self.outbound_closed && self.write_buffer.len() < MAX_STREAM_REF_FRAME_BYTES {
828 match self
829 .endpoint
830 .try_next_outbound(STREAM_REF_OUTBOUND_BATCH_FRAMES, MAX_STREAM_REF_FRAME_BYTES)
831 {
832 StreamRefOutboundPoll::Ready(Ok(outbound)) => {
833 encode_carrier_outbound_into(&outbound, &mut self.encode_buffer)?;
834 let encoded_len = self.encode_buffer.len();
835 if encoded_len == 0 {
836 continue;
837 }
838 if self.diagnostics.is_some() {
839 self.pending_diagnostics.push_back(PendingDiagnostic {
840 remaining: encoded_len,
841 counts: outbound_counts(&outbound),
842 });
843 }
844 self.outbound_written = true;
845 self.write_buffer.extend_from_slice(&self.encode_buffer);
846 }
847 StreamRefOutboundPoll::Ready(Err(error)) => return Err(error),
848 StreamRefOutboundPoll::Pending => {
849 self.recheck_outbound =
850 !self.outbound_written || self.inbound_seen || self.read_closed;
851 break;
852 }
853 StreamRefOutboundPoll::Closed => {
854 self.outbound_closed = true;
855 break;
856 }
857 }
858 }
859 Ok(())
860 }
861
862 fn handle_written(&mut self, written: Result<usize, std::io::Error>) -> StreamResult<()> {
863 match written {
864 Ok(0) => Err(StreamError::Failed(
865 "StreamRefs QUIC stream accepted zero write bytes".to_owned(),
866 )),
867 Ok(written) => {
868 self.write_buffer.advance(written);
869 self.record_written_bytes(written);
870 Ok(())
871 }
872 Err(error) => Err(io_error(error)),
873 }
874 }
875
876 fn record_written_bytes(&mut self, mut written: usize) {
877 let Some(diagnostics) = &self.diagnostics else {
878 return;
879 };
880 while written > 0 {
881 let Some(front) = self.pending_diagnostics.front_mut() else {
882 return;
883 };
884 if written < front.remaining {
885 front.remaining -= written;
886 return;
887 }
888 written -= front.remaining;
889 let counts = front.counts;
890 self.pending_diagnostics.pop_front();
891 diagnostics.record_counts(counts);
892 }
893 }
894
895 fn feed_read_buffer(&mut self, read: usize) -> StreamResult<()> {
896 feed_read_bytes(
897 &mut self.decoder,
898 &self.endpoint,
899 self.read_mode,
900 &mut self.pending_tail,
901 &self.read_buffer[..read],
902 )?;
903 self.inbound_seen = true;
904 Ok(())
905 }
906
907 fn handle_eof(&mut self) -> StreamResult<()> {
908 if !self.pending_tail.is_empty() {
909 feed_inbound_chunk(&mut self.decoder, &self.endpoint, &self.pending_tail)?;
910 self.pending_tail.clear();
911 }
912 if self.read_mode.fail_on_eof {
913 self.endpoint
914 .fail_connection(StreamError::AbruptTermination);
915 }
916 self.read_closed = true;
917 self.recheck_outbound = true;
918 Ok(())
919 }
920}
921
922#[cfg(feature = "tcp")]
923async fn run_stream_ref_endpoint_tcp_task<E>(
924 stream: TcpStream,
925 endpoint: E,
926 diagnostics: Option<StreamRefProtocolDiagnostics>,
927) -> StreamResult<NotUsed>
928where
929 E: StreamRefProtoEndpointWake,
930{
931 let (wake_sender, wake_receiver) = tokio_mpsc::channel(1);
932 endpoint.install_outbound_wake(wake_sender.clone());
933 let _ = wake_sender.try_send(());
934
935 let result = TcpEndpointTask {
936 stream,
937 endpoint: endpoint.clone(),
938 diagnostics,
939 read_mode: CarrierReadMode::new(STREAM_REF_TCP_CHUNK_SIZE, true, true),
940 decoder: FrameDecoder::default(),
941 read_buffer: BytesMut::with_capacity(STREAM_REF_TCP_CHUNK_SIZE),
942 pending_tail: Vec::with_capacity(STREAM_REF_TCP_CHUNK_SIZE),
943 write_buffer: BytesMut::with_capacity(STREAM_REF_TCP_CHUNK_SIZE),
944 encode_buffer: Vec::with_capacity(STREAM_REF_TCP_CHUNK_SIZE),
945 pending_diagnostics: VecDeque::new(),
946 outbound_closed: false,
947 write_shutdown: false,
948 wake_receiver,
949 }
950 .run()
951 .await;
952
953 endpoint.clear_outbound_wake();
954 if let Err(error) = &result {
955 endpoint.fail_connection(error.clone());
956 }
957 result
958}
959
960#[cfg(feature = "tcp")]
961struct TcpEndpointTask<E>
962where
963 E: StreamRefProtoEndpointWake,
964{
965 stream: TcpStream,
966 endpoint: E,
967 diagnostics: Option<StreamRefProtocolDiagnostics>,
968 read_mode: CarrierReadMode,
969 decoder: FrameDecoder,
970 read_buffer: BytesMut,
971 pending_tail: Vec<u8>,
972 write_buffer: BytesMut,
973 encode_buffer: Vec<u8>,
974 pending_diagnostics: VecDeque<PendingDiagnostic>,
975 outbound_closed: bool,
976 write_shutdown: bool,
977 wake_receiver: tokio_mpsc::Receiver<()>,
978}
979
980#[cfg(feature = "tcp")]
981impl<E> TcpEndpointTask<E>
982where
983 E: StreamRefProtoEndpointWake,
984{
985 async fn run(mut self) -> StreamResult<NotUsed> {
986 self.stream.set_nodelay(true).map_err(io_error)?;
987 loop {
988 self.drain_outbound()?;
989 if !self.write_buffer.is_empty() || (self.outbound_closed && !self.write_shutdown) {
990 self.flush_write_buffer().await?;
991 }
992
993 tokio::select! {
994 biased;
995 wake = self.wake_receiver.recv() => {
996 if wake.is_none() && !self.outbound_closed {
997 self.drain_outbound()?;
998 }
999 }
1000 ready = self.stream.readable() => {
1001 ready.map_err(io_error)?;
1002 if self.read_available()? {
1003 return Ok(NotUsed);
1004 }
1005 }
1006 ready = self.stream.writable(), if !self.write_buffer.is_empty() || (self.outbound_closed && !self.write_shutdown) => {
1007 ready.map_err(io_error)?;
1008 self.flush_ready_write_buffer()?;
1009 }
1010 }
1011 }
1012 }
1013
1014 fn drain_outbound(&mut self) -> StreamResult<()> {
1015 while !self.outbound_closed && self.write_buffer.len() < MAX_STREAM_REF_FRAME_BYTES {
1016 match self
1017 .endpoint
1018 .try_next_outbound(STREAM_REF_OUTBOUND_BATCH_FRAMES, MAX_STREAM_REF_FRAME_BYTES)
1019 {
1020 StreamRefOutboundPoll::Ready(Ok(outbound)) => {
1021 encode_carrier_outbound_into(&outbound, &mut self.encode_buffer)?;
1022 let encoded_len = self.encode_buffer.len();
1023 if encoded_len == 0 {
1024 continue;
1025 }
1026 if self.diagnostics.is_some() {
1027 self.pending_diagnostics.push_back(PendingDiagnostic {
1028 remaining: encoded_len,
1029 counts: outbound_counts(&outbound),
1030 });
1031 }
1032 self.write_buffer.extend_from_slice(&self.encode_buffer);
1033 }
1034 StreamRefOutboundPoll::Ready(Err(error)) => return Err(error),
1035 StreamRefOutboundPoll::Pending => break,
1036 StreamRefOutboundPoll::Closed => {
1037 self.outbound_closed = true;
1038 break;
1039 }
1040 }
1041 }
1042 Ok(())
1043 }
1044
1045 async fn flush_write_buffer(&mut self) -> StreamResult<()> {
1046 if !self.write_buffer.is_empty() {
1047 self.stream.writable().await.map_err(io_error)?;
1048 self.flush_ready_write_buffer()?;
1049 }
1050 if self.outbound_closed && self.write_buffer.is_empty() && !self.write_shutdown {
1051 self.stream.shutdown().await.map_err(io_error)?;
1052 self.write_shutdown = true;
1053 }
1054 Ok(())
1055 }
1056
1057 fn flush_ready_write_buffer(&mut self) -> StreamResult<()> {
1058 while !self.write_buffer.is_empty() {
1059 match self.stream.try_write(&self.write_buffer) {
1060 Ok(0) => {
1061 return Err(StreamError::Failed(
1062 "StreamRefs TCP socket accepted zero write bytes".to_owned(),
1063 ));
1064 }
1065 Ok(written) => {
1066 self.write_buffer.advance(written);
1067 self.record_written_bytes(written);
1068 }
1069 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => return Ok(()),
1070 Err(error) => return Err(io_error(error)),
1071 }
1072 }
1073
1074 Ok(())
1075 }
1076
1077 fn record_written_bytes(&mut self, mut written: usize) {
1078 let Some(diagnostics) = &self.diagnostics else {
1079 return;
1080 };
1081 while written > 0 {
1082 let Some(front) = self.pending_diagnostics.front_mut() else {
1083 return;
1084 };
1085 if written < front.remaining {
1086 front.remaining -= written;
1087 return;
1088 }
1089 written -= front.remaining;
1090 let counts = front.counts;
1091 self.pending_diagnostics.pop_front();
1092 diagnostics.record_counts(counts);
1093 }
1094 }
1095
1096 fn read_available(&mut self) -> StreamResult<bool> {
1097 loop {
1098 self.read_buffer.reserve(self.read_mode.chunk_size);
1099 match self.stream.try_read_buf(&mut self.read_buffer) {
1100 Ok(0) => return self.handle_eof(),
1101 Ok(_) => {
1102 feed_read_bytes(
1103 &mut self.decoder,
1104 &self.endpoint,
1105 self.read_mode,
1106 &mut self.pending_tail,
1107 &self.read_buffer,
1108 )?;
1109 self.read_buffer.clear();
1110 self.drain_outbound()?;
1111 }
1112 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => return Ok(false),
1113 Err(error) => return Err(io_error(error)),
1114 }
1115 }
1116 }
1117
1118 fn handle_eof(&mut self) -> StreamResult<bool> {
1119 if !self.pending_tail.is_empty() {
1120 feed_inbound_chunk(&mut self.decoder, &self.endpoint, &self.pending_tail)?;
1121 self.pending_tail.clear();
1122 }
1123 if self.read_mode.fail_on_eof {
1124 self.endpoint
1125 .fail_connection(StreamError::AbruptTermination);
1126 }
1127 Ok(true)
1128 }
1129}
1130
1131fn feed_read_bytes<E>(
1132 decoder: &mut FrameDecoder,
1133 endpoint: &E,
1134 read_mode: CarrierReadMode,
1135 pending_tail: &mut Vec<u8>,
1136 read_buffer: &[u8],
1137) -> StreamResult<()>
1138where
1139 E: StreamRefProtoEndpoint,
1140{
1141 if read_mode.emit_available {
1142 if !pending_tail.is_empty() {
1143 pending_tail.extend_from_slice(read_buffer);
1144 feed_inbound_chunk(decoder, endpoint, pending_tail)?;
1145 pending_tail.clear();
1146 return Ok(());
1147 }
1148 return feed_inbound_chunk(decoder, endpoint, read_buffer);
1149 }
1150
1151 let mut offset = 0;
1152 if !pending_tail.is_empty() {
1153 let needed = read_mode.chunk_size - pending_tail.len();
1154 let take = needed.min(read_buffer.len());
1155 pending_tail.extend_from_slice(&read_buffer[..take]);
1156 offset += take;
1157 if pending_tail.len() == read_mode.chunk_size {
1158 feed_inbound_chunk(decoder, endpoint, pending_tail)?;
1159 pending_tail.clear();
1160 }
1161 }
1162
1163 while offset + read_mode.chunk_size <= read_buffer.len() {
1164 let next = offset + read_mode.chunk_size;
1165 feed_inbound_chunk(decoder, endpoint, &read_buffer[offset..next])?;
1166 offset = next;
1167 }
1168
1169 if offset < read_buffer.len() {
1170 pending_tail.extend_from_slice(&read_buffer[offset..]);
1171 }
1172 Ok(())
1173}
1174
1175fn feed_inbound_chunk<E>(decoder: &mut FrameDecoder, endpoint: &E, chunk: &[u8]) -> StreamResult<()>
1176where
1177 E: StreamRefProtoEndpoint,
1178{
1179 decoder.push_chunk(chunk, endpoint)
1180}
1181
1182#[cfg(feature = "tcp")]
1183fn bind_tcp_listener<A>(addr: A) -> StreamResult<(TcpListener, StreamRefTcpBinding, Handle)>
1184where
1185 A: ToSocketAddrs + Send + 'static,
1186{
1187 let runtime = stream_ref_tcp_runtime()?;
1188 let listener = runtime
1189 .block_on(async { TcpListener::bind(addr).await })
1190 .map_err(io_error)?;
1191 let local_addr = listener.local_addr().map_err(io_error)?;
1192 Ok((
1193 listener,
1194 StreamRefTcpBinding { local_addr },
1195 runtime.handle().clone(),
1196 ))
1197}
1198
1199#[cfg(feature = "tcp")]
1200fn connect_tcp_stream<A>(addr: A) -> StreamResult<(TcpStream, Handle)>
1201where
1202 A: ToSocketAddrs + Send + 'static,
1203{
1204 let runtime = stream_ref_tcp_runtime()?;
1205 let stream = runtime
1206 .block_on(async { TcpStream::connect(addr).await })
1207 .map_err(io_error)?;
1208 stream.set_nodelay(true).map_err(io_error)?;
1209 Ok((stream, runtime.handle().clone()))
1210}
1211
1212#[cfg(feature = "tcp")]
1213fn stream_ref_tcp_runtime() -> StreamResult<&'static Runtime> {
1214 static RUNTIME: OnceLock<Result<Runtime, String>> = OnceLock::new();
1215 match RUNTIME.get_or_init(|| {
1216 tokio::runtime::Builder::new_multi_thread()
1217 .thread_name("datum-streamref-tcp")
1218 .enable_all()
1219 .build()
1220 .map_err(|error| error.to_string())
1221 }) {
1222 Ok(runtime) => Ok(runtime),
1223 Err(error) => Err(StreamError::Failed(format!(
1224 "failed to start StreamRefs TCP runtime: {error}"
1225 ))),
1226 }
1227}
1228
1229#[cfg(feature = "tcp")]
1230fn current_tokio_handle() -> StreamResult<Handle> {
1231 Handle::try_current().map_err(|error| {
1232 StreamError::Failed(format!(
1233 "StreamRefs TCP stream helper requires a current Tokio runtime: {error}"
1234 ))
1235 })
1236}
1237
1238fn io_error(error: std::io::Error) -> StreamError {
1239 StreamError::Failed(error.to_string())
1240}
1241
1242#[cfg(feature = "quic")]
1243fn is_quic_teardown_loss(error: &StreamError) -> bool {
1244 matches!(error, StreamError::Failed(message) if message == "connection lost")
1245}
1246
1247fn spawn_endpoint_task<F>(handle: &Handle, run: F) -> EndpointTaskCompletion
1248where
1249 F: Future<Output = StreamResult<NotUsed>> + Send + 'static,
1250{
1251 let (sender, receiver) = mpsc::channel();
1252 let task = handle.spawn(async move {
1253 let result = run.await;
1254 let _ = sender.send(result);
1255 });
1256 EndpointTaskCompletion {
1257 receiver,
1258 task: Some(task),
1259 }
1260}
1261
1262fn encode_carrier_outbound_into(
1263 outbound: &StreamRefOutbound,
1264 bytes: &mut Vec<u8>,
1265) -> StreamResult<()> {
1266 bytes.clear();
1267 match outbound {
1268 StreamRefOutbound::Frame(frame) => append_protobuf_carrier_frame(frame, bytes)?,
1269 StreamRefOutbound::SequencedBatch(batch) => {
1270 append_compact_payload_batch(batch, bytes)?;
1271 }
1272 }
1273 Ok(())
1274}
1275
1276#[cfg(test)]
1277fn encode_carrier_frames(frames: &[StreamRefFrame]) -> StreamResult<Vec<u8>> {
1278 let mut bytes = Vec::new();
1279 let mut index = 0;
1280 while index < frames.len() {
1281 if sequenced_on_next(&frames[index]).is_some() {
1282 let end = sequenced_run_end(frames, index);
1283 append_compact_sequenced_batches(&frames[index..end], &mut bytes)?;
1284 index = end;
1285 } else {
1286 append_protobuf_carrier_frame(&frames[index], &mut bytes)?;
1287 index += 1;
1288 }
1289 }
1290 Ok(bytes)
1291}
1292
1293fn append_compact_payload_batch(
1294 batch: &StreamRefPayloadBatch,
1295 bytes: &mut Vec<u8>,
1296) -> StreamResult<()> {
1297 let mut start = 0;
1298 while start < batch.count() {
1299 let mut end = start;
1300 let mut payload_len = COMPACT_BATCH_HEADER_BYTES;
1301 while end < batch.count() {
1302 let element_len = COMPACT_BATCH_ELEMENT_LEN_BYTES
1303 .checked_add(batch.payload_len(end))
1304 .ok_or(StreamError::LimitExceeded {
1305 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1306 })?;
1307 let next_payload_len =
1308 payload_len
1309 .checked_add(element_len)
1310 .ok_or(StreamError::LimitExceeded {
1311 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1312 })?;
1313 if end > start
1314 && (next_payload_len > MAX_STREAM_REF_FRAME_BYTES
1315 || end - start >= u16::MAX as usize)
1316 {
1317 break;
1318 }
1319 if next_payload_len > MAX_STREAM_REF_FRAME_BYTES {
1320 return Err(StreamError::LimitExceeded {
1321 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1322 });
1323 }
1324 payload_len = next_payload_len;
1325 end += 1;
1326 }
1327 append_compact_payload_batch_slice(batch, start, end, payload_len, bytes)?;
1328 start = end;
1329 }
1330 Ok(())
1331}
1332
1333fn append_compact_payload_batch_slice(
1334 batch: &StreamRefPayloadBatch,
1335 start: usize,
1336 end: usize,
1337 payload_len: usize,
1338 bytes: &mut Vec<u8>,
1339) -> StreamResult<()> {
1340 let payload_len = u32::try_from(payload_len).map_err(|_| StreamError::LimitExceeded {
1341 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1342 })?;
1343 let count = u16::try_from(end - start).map_err(|_| StreamError::LimitExceeded {
1344 max: u16::MAX as u64,
1345 })?;
1346 let first_seq = batch
1347 .first_seq_nr()
1348 .checked_add(start as u64)
1349 .ok_or_else(|| StreamError::Failed("compact StreamRefs seq_nr overflow".to_owned()))?;
1350 bytes.extend((COMPACT_FRAME_FLAG | payload_len).to_be_bytes());
1351 bytes.push(COMPACT_FRAME_VERSION);
1352 bytes.push(COMPACT_SEQUENCED_ON_NEXT_BATCH);
1353 bytes.extend(batch.stream_ref_id().to_bytes());
1354 bytes.extend(first_seq.to_be_bytes());
1355 bytes.extend(count.to_be_bytes());
1356 for index in start..end {
1357 let payload = batch.payload(index);
1358 let payload_len = u32::try_from(payload.len()).map_err(|_| StreamError::LimitExceeded {
1359 max: u32::MAX as u64,
1360 })?;
1361 bytes.extend(payload_len.to_be_bytes());
1362 bytes.extend(payload);
1363 }
1364 Ok(())
1365}
1366
1367fn append_protobuf_carrier_frame(frame: &StreamRefFrame, bytes: &mut Vec<u8>) -> StreamResult<()> {
1368 let payload = frame.encode_to_vec();
1369 let len = u32::try_from(payload.len()).map_err(|_| StreamError::LimitExceeded {
1370 max: COMPACT_FRAME_LEN_MASK as u64,
1371 })?;
1372 if payload.len() > MAX_STREAM_REF_FRAME_BYTES || len > COMPACT_FRAME_LEN_MASK {
1373 return Err(StreamError::LimitExceeded {
1374 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1375 });
1376 }
1377 bytes.extend(len.to_be_bytes());
1378 bytes.extend(payload);
1379 Ok(())
1380}
1381
1382#[cfg(test)]
1383fn append_compact_sequenced_batches(
1384 frames: &[StreamRefFrame],
1385 bytes: &mut Vec<u8>,
1386) -> StreamResult<()> {
1387 let mut start = 0;
1388 while start < frames.len() {
1389 let mut end = start;
1390 let mut payload_len = COMPACT_BATCH_HEADER_BYTES;
1391 while end < frames.len() {
1392 let (_, payload) = sequenced_on_next(&frames[end]).expect("sequenced frame");
1393 let element_len = COMPACT_BATCH_ELEMENT_LEN_BYTES
1394 .checked_add(payload.len())
1395 .ok_or(StreamError::LimitExceeded {
1396 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1397 })?;
1398 let next_payload_len =
1399 payload_len
1400 .checked_add(element_len)
1401 .ok_or(StreamError::LimitExceeded {
1402 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1403 })?;
1404 if end > start
1405 && (next_payload_len > MAX_STREAM_REF_FRAME_BYTES
1406 || end - start >= u16::MAX as usize)
1407 {
1408 break;
1409 }
1410 if next_payload_len > MAX_STREAM_REF_FRAME_BYTES {
1411 return Err(StreamError::LimitExceeded {
1412 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1413 });
1414 }
1415 payload_len = next_payload_len;
1416 end += 1;
1417 }
1418 append_compact_sequenced_batch(&frames[start..end], payload_len, bytes)?;
1419 start = end;
1420 }
1421 Ok(())
1422}
1423
1424#[cfg(test)]
1425fn append_compact_sequenced_batch(
1426 frames: &[StreamRefFrame],
1427 payload_len: usize,
1428 bytes: &mut Vec<u8>,
1429) -> StreamResult<()> {
1430 let (first_seq, _) = sequenced_on_next(&frames[0]).expect("sequenced frame");
1431 let payload_len = u32::try_from(payload_len).map_err(|_| StreamError::LimitExceeded {
1432 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1433 })?;
1434 let count = u16::try_from(frames.len()).map_err(|_| StreamError::LimitExceeded {
1435 max: u16::MAX as u64,
1436 })?;
1437 bytes.extend((COMPACT_FRAME_FLAG | payload_len).to_be_bytes());
1438 bytes.push(COMPACT_FRAME_VERSION);
1439 bytes.push(COMPACT_SEQUENCED_ON_NEXT_BATCH);
1440 bytes.extend(frames[0].stream_ref_id.to_bytes());
1441 bytes.extend(first_seq.to_be_bytes());
1442 bytes.extend(count.to_be_bytes());
1443 for frame in frames {
1444 let (_, payload) = sequenced_on_next(frame).expect("sequenced frame");
1445 let payload_len = u32::try_from(payload.len()).map_err(|_| StreamError::LimitExceeded {
1446 max: u32::MAX as u64,
1447 })?;
1448 bytes.extend(payload_len.to_be_bytes());
1449 bytes.extend(payload);
1450 }
1451 Ok(())
1452}
1453
1454#[cfg(test)]
1455fn sequenced_run_end(frames: &[StreamRefFrame], start: usize) -> usize {
1456 let mut end = start + 1;
1457 while end < frames.len() {
1458 let Some((previous_seq, _)) = sequenced_on_next(&frames[end - 1]) else {
1459 break;
1460 };
1461 let Some((next_seq, _)) = sequenced_on_next(&frames[end]) else {
1462 break;
1463 };
1464 if frames[end].stream_ref_id != frames[start].stream_ref_id
1465 || next_seq != previous_seq.saturating_add(1)
1466 {
1467 break;
1468 }
1469 end += 1;
1470 }
1471 end
1472}
1473
1474#[cfg(test)]
1475fn sequenced_on_next(frame: &StreamRefFrame) -> Option<(u64, &[u8])> {
1476 match &frame.message {
1477 StreamRefMessage::SequencedOnNext { seq_nr, payload } => {
1478 Some((*seq_nr, payload.bytes.as_slice()))
1479 }
1480 _ => None,
1481 }
1482}
1483
1484#[derive(Default)]
1485struct FrameDecoder {
1486 buffer: BytesMut,
1487 offset: usize,
1488}
1489
1490impl FrameDecoder {
1491 fn push_chunk<E>(&mut self, chunk: &[u8], endpoint: &E) -> StreamResult<()>
1492 where
1493 E: StreamRefProtoEndpoint,
1494 {
1495 self.buffer.extend_from_slice(chunk);
1496 while let Some(header) = self.peek_header()? {
1497 if self.buffer.len().saturating_sub(self.offset) < FRAME_LEN_BYTES + header.len {
1498 break;
1499 }
1500 let payload_start = self.offset + FRAME_LEN_BYTES;
1501 let payload_end = payload_start + header.len;
1502 let payload = &self.buffer[payload_start..payload_end];
1503 match header.kind {
1504 CarrierFrameKind::Protobuf => {
1505 endpoint.handle_frame(StreamRefFrame::decode(payload)?)?;
1506 }
1507 CarrierFrameKind::Compact => {
1508 decode_compact_carrier_frame(payload, endpoint)?;
1509 }
1510 }
1511 self.offset = payload_end;
1512 }
1513 if self.offset > 0 && (self.offset == self.buffer.len() || self.offset >= 64 * 1024) {
1514 self.buffer.advance(self.offset);
1515 self.offset = 0;
1516 }
1517 Ok(())
1518 }
1519
1520 fn peek_header(&self) -> StreamResult<Option<CarrierFrameHeader>> {
1521 if self.buffer.len().saturating_sub(self.offset) < FRAME_LEN_BYTES {
1522 return Ok(None);
1523 }
1524 let len = self.buffer[self.offset..self.offset + FRAME_LEN_BYTES]
1525 .try_into()
1526 .expect("frame header length");
1527 let raw_len = u32::from_be_bytes(len);
1528 let kind = if raw_len & COMPACT_FRAME_FLAG == 0 {
1529 CarrierFrameKind::Protobuf
1530 } else {
1531 CarrierFrameKind::Compact
1532 };
1533 let len = (raw_len & COMPACT_FRAME_LEN_MASK) as usize;
1534 if len > MAX_STREAM_REF_FRAME_BYTES {
1535 return Err(StreamError::LimitExceeded {
1536 max: MAX_STREAM_REF_FRAME_BYTES as u64,
1537 });
1538 }
1539 Ok(Some(CarrierFrameHeader { kind, len }))
1540 }
1541}
1542
1543#[derive(Clone, Copy)]
1544struct CarrierFrameHeader {
1545 kind: CarrierFrameKind,
1546 len: usize,
1547}
1548
1549#[derive(Clone, Copy)]
1550enum CarrierFrameKind {
1551 Protobuf,
1552 Compact,
1553}
1554
1555fn decode_compact_carrier_frame<E>(payload: &[u8], endpoint: &E) -> StreamResult<()>
1556where
1557 E: StreamRefProtoEndpoint,
1558{
1559 if payload.len() < COMPACT_BATCH_HEADER_BYTES {
1560 return Err(StreamError::Failed(
1561 "compact StreamRefs carrier frame too short".to_owned(),
1562 ));
1563 }
1564 let version = payload[0];
1565 if version != COMPACT_FRAME_VERSION {
1566 return Err(StreamError::Failed(format!(
1567 "unsupported compact StreamRefs carrier frame version: {version}"
1568 )));
1569 }
1570 let kind = payload[1];
1571 if kind != COMPACT_SEQUENCED_ON_NEXT_BATCH {
1572 return Err(StreamError::Failed(format!(
1573 "unsupported compact StreamRefs carrier frame kind: {kind}"
1574 )));
1575 }
1576 let stream_ref_id = StreamRefId::from_bytes(&payload[2..18])?;
1577 let first_seq = u64::from_be_bytes(payload[18..26].try_into().expect("seq len"));
1578 let count = u16::from_be_bytes(payload[26..28].try_into().expect("count len")) as usize;
1579 if count == 0 {
1580 return Err(StreamError::Failed(
1581 "compact StreamRefs carrier batch is empty".to_owned(),
1582 ));
1583 }
1584
1585 let mut offset = COMPACT_BATCH_HEADER_BYTES;
1586 let mut payloads = Vec::with_capacity(count);
1587 for index in 0..count {
1588 if payload.len().saturating_sub(offset) < COMPACT_BATCH_ELEMENT_LEN_BYTES {
1589 return Err(StreamError::Failed(
1590 "compact StreamRefs carrier batch has truncated payload length".to_owned(),
1591 ));
1592 }
1593 let payload_len = u32::from_be_bytes(
1594 payload[offset..offset + COMPACT_BATCH_ELEMENT_LEN_BYTES]
1595 .try_into()
1596 .expect("payload len"),
1597 ) as usize;
1598 offset += COMPACT_BATCH_ELEMENT_LEN_BYTES;
1599 if payload.len().saturating_sub(offset) < payload_len {
1600 return Err(StreamError::Failed(
1601 "compact StreamRefs carrier batch has truncated payload".to_owned(),
1602 ));
1603 }
1604 first_seq
1605 .checked_add(index as u64)
1606 .ok_or_else(|| StreamError::Failed("compact StreamRefs seq_nr overflow".to_owned()))?;
1607 payloads.push(&payload[offset..offset + payload_len]);
1608 offset += payload_len;
1609 }
1610 if offset != payload.len() {
1611 return Err(StreamError::Failed(
1612 "compact StreamRefs carrier batch has trailing bytes".to_owned(),
1613 ));
1614 }
1615 endpoint.handle_sequenced_on_next_batch(stream_ref_id, first_seq, &payloads)
1616}
1617
1618#[cfg(test)]
1619mod tests {
1620 use super::*;
1621 use std::sync::{Arc, Mutex};
1622
1623 #[derive(Clone)]
1624 struct RecordingEndpoint {
1625 stream_ref_id: StreamRefId,
1626 frames: Arc<Mutex<Vec<StreamRefFrame>>>,
1627 }
1628
1629 impl RecordingEndpoint {
1630 fn new(stream_ref_id: StreamRefId) -> Self {
1631 Self {
1632 stream_ref_id,
1633 frames: Arc::new(Mutex::new(Vec::new())),
1634 }
1635 }
1636
1637 fn frames(&self) -> Vec<StreamRefFrame> {
1638 self.frames.lock().expect("recording endpoint").clone()
1639 }
1640 }
1641
1642 impl StreamRefProtoEndpoint for RecordingEndpoint {
1643 fn stream_ref_id(&self) -> StreamRefId {
1644 self.stream_ref_id
1645 }
1646
1647 fn next_frame(&self) -> Option<StreamResult<StreamRefFrame>> {
1648 None
1649 }
1650
1651 fn handle_frame(&self, frame: StreamRefFrame) -> StreamResult<()> {
1652 self.frames.lock().expect("recording endpoint").push(frame);
1653 Ok(())
1654 }
1655
1656 fn fail_connection(&self, _error: StreamError) {}
1657 }
1658
1659 #[test]
1660 fn carrier_frame_decoder_reassembles_split_frames() {
1661 let frame = StreamRefFrame::new(
1662 StreamRefId::from_u128(1),
1663 datum::StreamRefMessage::CumulativeDemand { seq_nr: 32 },
1664 );
1665 let bytes = encode_carrier_frames(std::slice::from_ref(&frame)).unwrap();
1666 let split = bytes.len() / 2;
1667 let mut decoder = FrameDecoder::default();
1668 let endpoint = RecordingEndpoint::new(StreamRefId::from_u128(1));
1669
1670 decoder.push_chunk(&bytes[..split], &endpoint).unwrap();
1671 assert!(endpoint.frames().is_empty());
1672 decoder.push_chunk(&bytes[split..], &endpoint).unwrap();
1673 assert_eq!(endpoint.frames(), vec![frame]);
1674 }
1675
1676 #[test]
1677 fn compact_carrier_batch_round_trips_sequenced_frames() {
1678 let frames = (0_u64..3)
1679 .map(|seq_nr| {
1680 StreamRefFrame::new(
1681 StreamRefId::from_u128(7),
1682 datum::StreamRefMessage::SequencedOnNext {
1683 seq_nr,
1684 payload: datum::StreamRefPayloadBytes {
1685 bytes: seq_nr.to_be_bytes().to_vec(),
1686 },
1687 },
1688 )
1689 })
1690 .collect::<Vec<_>>();
1691 let bytes = encode_carrier_frames(&frames).unwrap();
1692 let header = u32::from_be_bytes(bytes[..FRAME_LEN_BYTES].try_into().unwrap());
1693 assert_ne!(header & COMPACT_FRAME_FLAG, 0);
1694
1695 let mut decoder = FrameDecoder::default();
1696 let endpoint = RecordingEndpoint::new(StreamRefId::from_u128(7));
1697 decoder.push_chunk(&bytes, &endpoint).unwrap();
1698 assert_eq!(endpoint.frames(), frames);
1699 }
1700
1701 #[test]
1702 fn compact_carrier_batch_reassembles_split_frames() {
1703 let frames = (4_u64..8)
1704 .map(|seq_nr| {
1705 StreamRefFrame::new(
1706 StreamRefId::from_u128(8),
1707 datum::StreamRefMessage::SequencedOnNext {
1708 seq_nr,
1709 payload: datum::StreamRefPayloadBytes {
1710 bytes: vec![seq_nr as u8],
1711 },
1712 },
1713 )
1714 })
1715 .collect::<Vec<_>>();
1716 let bytes = encode_carrier_frames(&frames).unwrap();
1717 let split = FRAME_LEN_BYTES + 5;
1718 let mut decoder = FrameDecoder::default();
1719 let endpoint = RecordingEndpoint::new(StreamRefId::from_u128(8));
1720
1721 decoder.push_chunk(&bytes[..split], &endpoint).unwrap();
1722 assert!(endpoint.frames().is_empty());
1723 decoder.push_chunk(&bytes[split..], &endpoint).unwrap();
1724 assert_eq!(endpoint.frames(), frames);
1725 }
1726}