Skip to main content

datum_net/
stream_ref.rs

1//! Remote StreamRefs carriers.
2//!
3//! `datum-core` owns the protobuf protocol and state machine. This module is
4//! the carrier layer: it length-prefixes protobuf frames and pumps them over
5//! reliable, ordered bidirectional byte streams such as QUIC or plaintext TCP.
6
7#[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
46// Carrier wire v1 keeps protobuf for control frames. A high bit in the
47// length-prefix marks compact SequencedOnNext batches: version, kind,
48// stream_ref_id, first seqNr, count, then length-prefixed payloads.
49const 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/// Counts selected StreamRefs protocol messages successfully written by a
75/// carrier endpoint.
76#[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/// Shared collector for StreamRefs protocol message-count diagnostics.
84///
85/// The carrier records counts only after the encoded frame has been written
86/// successfully. This is intentionally transport-level instrumentation; it does
87/// not change the protobuf protocol or wire format.
88#[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/// Completion handle for a StreamRefs-over-QUIC carrier.
156#[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/// Local TCP listener binding used by StreamRefs-over-TCP producer endpoints.
180#[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/// Completion handle for a StreamRefs-over-TCP carrier.
195#[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/// Serves a local `SourceRef` over an accepted or opened QUIC bidi stream.
251#[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/// Serves a local `Source` over an accepted or opened QUIC bidi stream.
266#[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/// Creates a local source fed by a remote QUIC StreamRef producer.
282#[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/// Serves a local `SinkRef` receiver over an accepted or opened QUIC bidi
298/// stream, returning a [`Source`] of inbound elements.
299///
300/// This is the local/receiver side of the SinkRef-over-QUIC pair: the remote
301/// sender pushes elements into a [`sink_ref_over_quic`](fn.sink_ref_over_quic)
302/// `Sink`, and this side surfaces them as a `Source`. The caller runs the
303/// returned source into a local `Sink` (for example `Sink::collect` or a fold).
304#[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/// Creates a local `Sink` that sends its incoming elements over QUIC to a
320/// remote `SinkRef` receiver.
321#[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/// Serves a local `SourceRef` over a one-shot plaintext TCP listener.
337///
338/// The listener binds immediately, accepts one connection, then runs the same
339/// StreamRefs protocol used by the QUIC carrier. The remote receiver should use
340/// [`source_ref_over_tcp`] to open the TCP connection and send the initial
341/// subscribe+demand frames.
342#[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/// Serves a local `SourceRef` over an already-connected Tokio TCP stream.
377///
378/// This is the stream-shaped counterpart to [`serve_source_ref_over_tcp`],
379/// intended for callers that own listener lifecycle separately, such as a
380/// benchmark server that reuses one bound listener across operations.
381#[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/// Creates a local source fed by a remote plaintext TCP StreamRef producer.
422///
423/// This receiver side opens the TCP connection and, once the returned source is
424/// materialized, sends the initial handshake and cumulative demand.
425#[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/// Creates a local source fed by a remote producer over an already-connected
457/// Tokio TCP stream.
458#[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/// Serves a remote `SinkRef` receiver over plaintext TCP.
488///
489/// This receiver side opens the TCP connection to a sender created with
490/// [`sink_ref_over_tcp`] and, once the returned source is materialized, sends
491/// the initial handshake and cumulative demand.
492#[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/// Serves a remote `SinkRef` receiver over an already-connected Tokio TCP
524/// stream.
525#[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/// Creates a local `Sink` whose incoming elements are sent over a one-shot
555/// plaintext TCP listener to a remote `SinkRef` receiver.
556///
557/// The producer/listener side waits for the receiver to open the TCP
558/// connection. This mirrors the QUIC SinkRef direction where the receiver
559/// opens the bidi stream so its handshake+demand establish the stream before
560/// the producer has elements to send.
561#[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/// Creates a local `Sink` whose incoming elements are sent over an
602/// already-connected Tokio TCP stream to a remote `SinkRef` receiver.
603#[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}