Skip to main content

microsandbox_agentd/
tcp.rs

1//! Guest-side TCP stream session handling.
2//!
3//! Handles `core.tcp.*` protocol messages by opening TCP sockets from
4//! inside the guest and relaying bytes between those sockets and the host.
5
6use std::time::Duration;
7
8use bytes::Bytes;
9use tokio::io::{AsyncReadExt, AsyncWrite, AsyncWriteExt};
10use tokio::net::TcpStream;
11use tokio::sync::{mpsc, watch};
12use tokio::task::JoinHandle;
13
14use microsandbox_protocol::bulk::{
15    BULK_FLOW_MASK_GUEST_TO_HOST, BULK_FLOW_MASK_HOST_TO_GUEST, BulkAccepted, BulkCredit,
16    BulkFinish, BulkFlow, BulkKind, BulkOffer, BulkReceiveState, BulkRecord, BulkSendState,
17    DEFAULT_BULK_RECORD_PAYLOAD, DEFAULT_BULK_WINDOW, MIN_BULK_RECORD_PAYLOAD,
18};
19use microsandbox_protocol::codec;
20use microsandbox_protocol::message::{Message, MessageType};
21use microsandbox_protocol::tcp::{TcpClosed, TcpConnect, TcpConnected, TcpData, TcpEof, TcpFailed};
22
23use crate::agent::{AdmittedBulkRecord, BulkInputPermit};
24use crate::serial::InputCharge;
25#[cfg(test)]
26use crate::session::SessionOutputEnvelope;
27use crate::session::{
28    BulkSessionOutput, RawActivity, RawSessionCompletion, RawSessionOutput, SessionOutput,
29    SessionOutputPermit, SessionOutputSender,
30};
31
32//--------------------------------------------------------------------------------------------------
33// Constants
34//--------------------------------------------------------------------------------------------------
35
36/// TCP stream read chunk size.
37const TCP_CHUNK_SIZE: usize = 64 * 1024;
38
39/// Capacity reserved before cloning and encoding one TCP data chunk.
40///
41/// The factor of two covers CBOR/framing overhead and allocator growth while keeping hundreds of
42/// TCP chunks eligible under the aggregate 32 MiB output budget.
43const TCP_OUTPUT_RESERVATION: usize = 2 * TCP_CHUNK_SIZE;
44
45/// Enough host-to-guest data slots for the default byte window at the smallest negotiated record.
46/// Lifecycle and absolute-credit updates use separate channels and cannot consume these slots.
47const TCP_COMMAND_CAPACITY: usize = DEFAULT_BULK_WINDOW as usize / MIN_BULK_RECORD_PAYLOAD as usize;
48
49/// Upper bound on a single guest-side connect attempt. The connect runs in the
50/// per-session task, so this only bounds that task's lifetime; it never blocks
51/// the agent's serial loop.
52const TCP_CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
53
54//--------------------------------------------------------------------------------------------------
55// Types
56//--------------------------------------------------------------------------------------------------
57
58/// Tracks an active guest-originated TCP stream.
59pub struct TcpSession {
60    owner_id: u32,
61    commands: mpsc::Sender<TcpCommand>,
62    bulk_control: Option<TcpBulkControlSenders>,
63    task: JoinHandle<()>,
64    bulk: bool,
65}
66
67enum TcpCommand {
68    Data(Vec<u8>, Option<InputCharge>),
69    Eof(Option<InputCharge>),
70    BulkRecord(AdmittedBulkRecord),
71}
72
73/// Coalescing lifecycle channels that cannot be starved by a full TCP data queue.
74struct TcpBulkControlSenders {
75    credit: watch::Sender<Option<BulkCredit>>,
76    finish: mpsc::Sender<BulkFinish>,
77}
78
79struct TcpBulkControlReceivers {
80    credit: watch::Receiver<Option<BulkCredit>>,
81    finish: mpsc::Receiver<BulkFinish>,
82}
83
84struct TcpBulkState {
85    send: BulkSendState,
86    receive: BulkReceiveState,
87}
88
89/// One ordered host-to-guest record being written incrementally. Keeping at most one record here
90/// preserves credit semantics while allowing the opposite TCP half to keep reading concurrently.
91struct PendingTcpWrite {
92    payload: Bytes,
93    written: usize,
94    bulk_end: Option<u64>,
95    _bulk_input_permit: Option<BulkInputPermit>,
96    _control_input_charge: Option<InputCharge>,
97}
98
99//--------------------------------------------------------------------------------------------------
100// Methods
101//--------------------------------------------------------------------------------------------------
102
103impl TcpSession {
104    /// Correlation ID whose relay client owns this TCP stream.
105    pub fn owner_id(&self) -> u32 {
106        self.owner_id
107    }
108
109    /// Queue stream data to write to the guest socket.
110    ///
111    /// Awaits queue space when the per-session relay is behind, so a stalled
112    /// destination backpressures the caller instead of growing memory.
113    pub async fn write_data(&self, data: Vec<u8>) -> Result<(), String> {
114        self.write_data_charged(data, None).await
115    }
116
117    pub(crate) async fn write_data_charged(
118        &self,
119        data: Vec<u8>,
120        charge: Option<InputCharge>,
121    ) -> Result<(), String> {
122        if self.bulk {
123            return Err("CBOR TCP data is invalid after raw bulk acceptance".into());
124        }
125        self.commands
126            .send(TcpCommand::Data(data, charge))
127            .await
128            .map_err(|_| "TCP session is closed".to_string())
129    }
130
131    /// Close the guest socket write half.
132    ///
133    /// Ordered after any queued data, so the destination sees the write shutdown
134    /// only once it has received everything sent before it.
135    pub async fn close_write(&self) -> Result<(), String> {
136        self.close_write_charged(None).await
137    }
138
139    pub(crate) async fn close_write_charged(
140        &self,
141        charge: Option<InputCharge>,
142    ) -> Result<(), String> {
143        if self.bulk {
144            return Err("CBOR TCP EOF is invalid after raw bulk acceptance".into());
145        }
146        self.commands
147            .send(TcpCommand::Eof(charge))
148            .await
149            .map_err(|_| "TCP session is closed".to_string())
150    }
151
152    /// Queue one host-to-guest raw bulk record.
153    pub(crate) async fn write_bulk(&self, record: AdmittedBulkRecord) -> Result<(), String> {
154        if !self.bulk {
155            return Err("raw bulk record sent to a generation-6 TCP stream".into());
156        }
157        self.commands
158            .try_send(TcpCommand::BulkRecord(record))
159            .map_err(|error| format!("TCP bulk input queue is unavailable: {error}"))
160    }
161
162    /// Deliver an absolute guest-to-host credit update.
163    pub async fn apply_credit(&self, credit: BulkCredit) -> Result<(), String> {
164        if !self.bulk {
165            return Err("bulk credit sent to a generation-6 TCP stream".into());
166        }
167        let control = self
168            .bulk_control
169            .as_ref()
170            .ok_or_else(|| "TCP bulk control path is unavailable".to_string())?;
171        if control.credit.is_closed() {
172            // Sink consumption can return credit after the producer queued its final output.
173            // There is no sender left to enable; failing here would cancel its queued raw tail.
174            return Ok(());
175        }
176        control.credit.send_replace(Some(credit));
177        Ok(())
178    }
179
180    /// Queue the exact host-to-guest half-close marker.
181    pub async fn finish_bulk(&self, finish: BulkFinish) -> Result<(), String> {
182        if !self.bulk {
183            return Err("bulk finish sent to a generation-6 TCP stream".into());
184        }
185        let control = self
186            .bulk_control
187            .as_ref()
188            .ok_or_else(|| "TCP bulk control path is unavailable".to_string())?;
189        control
190            .finish
191            .try_send(finish)
192            .map_err(|error| format!("TCP bulk finish path is unavailable: {error}"))
193    }
194
195    /// Whether this TCP session negotiated generation-8 raw bulk.
196    pub fn is_bulk(&self) -> bool {
197        self.bulk
198    }
199
200    /// Tear down the TCP session.
201    ///
202    /// Aborts the relay task directly rather than queuing a command, so teardown
203    /// never waits behind a full command queue. Dropping the task closes the
204    /// guest socket. The host has already closed its side before asking for this,
205    /// so no terminal frame is owed back to it.
206    pub fn close(&self) {
207        self.task.abort();
208    }
209
210    /// Returns whether the background relay task has finished.
211    pub fn is_finished(&self) -> bool {
212        self.task.is_finished()
213    }
214
215    /// Open a TCP stream from inside the guest and start relaying it.
216    ///
217    /// The OS connect runs inside the spawned task, not on the caller's serial
218    /// loop, so a hanging or slow destination can never wedge the agent. The
219    /// task reports `core.tcp.connected` on success or a terminal
220    /// `core.tcp.failed` on error/timeout over `session_tx`; the host correlates
221    /// either reply by id. The returned session is live immediately, with
222    /// commands queued until the connect completes.
223    pub fn open(id: u32, req: TcpConnect, session_tx: &SessionOutputSender) -> Self {
224        let bulk = req.bulk.is_some();
225        let (commands_tx, commands_rx) = mpsc::channel(TCP_COMMAND_CAPACITY);
226        let (bulk_control, bulk_control_rx) = if bulk {
227            let (credit_tx, credit_rx) = watch::channel(None);
228            let (finish_tx, finish_rx) = mpsc::channel(1);
229            (
230                Some(TcpBulkControlSenders {
231                    credit: credit_tx,
232                    finish: finish_tx,
233                }),
234                Some(TcpBulkControlReceivers {
235                    credit: credit_rx,
236                    finish: finish_rx,
237                }),
238            )
239        } else {
240            (None, None)
241        };
242        let output_tx = session_tx.clone();
243        let task = tokio::spawn(async move {
244            connect_and_relay(id, req, commands_rx, bulk_control_rx, output_tx).await;
245        });
246
247        Self {
248            owner_id: id,
249            commands: commands_tx,
250            bulk_control,
251            task,
252            bulk,
253        }
254    }
255}
256
257//--------------------------------------------------------------------------------------------------
258// Functions: Helpers
259//--------------------------------------------------------------------------------------------------
260
261/// Connects to the destination, reports the outcome, then relays the stream.
262///
263/// Runs entirely inside the per-session task. On a connect error or timeout it
264/// emits a terminal `core.tcp.failed`; the agent loop removes the session when
265/// that frame flows past. On success it emits `core.tcp.connected` and hands off
266/// to the relay loop.
267async fn connect_and_relay(
268    id: u32,
269    req: TcpConnect,
270    commands: mpsc::Receiver<TcpCommand>,
271    bulk_control: Option<TcpBulkControlReceivers>,
272    tx: SessionOutputSender,
273) {
274    let TcpConnect { host, port, bulk } = req;
275    let connect = TcpStream::connect((host.as_str(), port));
276    let stream = match tokio::time::timeout(TCP_CONNECT_TIMEOUT, connect).await {
277        Ok(Ok(stream)) => stream,
278        Ok(Err(e)) => {
279            send_raw_tcp_message(
280                id,
281                MessageType::TcpFailed,
282                &TcpFailed {
283                    error: format!("connect {host}:{port}: {e}"),
284                },
285                RawActivity::guest_message(),
286                Some(RawSessionCompletion::Tcp),
287                &tx,
288            )
289            .await;
290            return;
291        }
292        Err(_elapsed) => {
293            send_raw_tcp_message(
294                id,
295                MessageType::TcpFailed,
296                &TcpFailed {
297                    error: format!("connect {host}:{port} timed out"),
298                },
299                RawActivity::guest_message(),
300                Some(RawSessionCompletion::Tcp),
301                &tx,
302            )
303            .await;
304            return;
305        }
306    };
307
308    if !send_raw_tcp_message(
309        id,
310        MessageType::TcpConnected,
311        &TcpConnected {},
312        RawActivity::guest_message(),
313        None,
314        &tx,
315    )
316    .await
317    {
318        return;
319    }
320
321    let bulk = match bulk {
322        Some(offer) => {
323            let accepted = match accept_tcp_offer(offer) {
324                Ok(accepted) => accepted,
325                Err(error) => {
326                    send_raw_tcp_message(
327                        id,
328                        MessageType::TcpFailed,
329                        &TcpFailed { error },
330                        RawActivity::guest_message(),
331                        Some(RawSessionCompletion::Tcp),
332                        &tx,
333                    )
334                    .await;
335                    return;
336                }
337            };
338            if !send_raw_tcp_message(
339                id,
340                MessageType::BulkAccepted,
341                &accepted,
342                RawActivity::guest_message(),
343                None,
344                &tx,
345            )
346            .await
347            {
348                return;
349            }
350            let send = match BulkSendState::new(
351                BulkKind::Tcp,
352                BulkFlow::GuestToHost,
353                accepted.max_record_payload,
354                accepted.guest_to_host_credit_limit,
355            ) {
356                Ok(send) => send,
357                Err(error) => {
358                    eprintln!("failed to create TCP bulk send state for {id}: {error}");
359                    return;
360                }
361            };
362            let receive = match BulkReceiveState::new(
363                BulkKind::Tcp,
364                BulkFlow::HostToGuest,
365                accepted.max_record_payload,
366                accepted.host_to_guest_credit_limit,
367                DEFAULT_BULK_WINDOW,
368            ) {
369                Ok(receive) => receive,
370                Err(error) => {
371                    eprintln!("failed to create TCP bulk receive state for {id}: {error}");
372                    return;
373                }
374            };
375            Some(TcpBulkState { send, receive })
376        }
377        None => None,
378    };
379
380    relay_tcp_session(id, stream, commands, bulk_control, tx, bulk).await;
381}
382
383fn accept_tcp_offer(offer: BulkOffer) -> Result<BulkAccepted, String> {
384    let offer = offer
385        .validate()
386        .map_err(|error| format!("invalid TCP bulk offer: {error}"))?;
387    if offer.guest_to_host_credit_limit == 0 {
388        return Err("TCP bulk offer must grant guest-to-host credit".into());
389    }
390    Ok(BulkAccepted {
391        kind: BulkKind::Tcp,
392        flows: BULK_FLOW_MASK_HOST_TO_GUEST | BULK_FLOW_MASK_GUEST_TO_HOST,
393        format: offer.format,
394        max_record_payload: offer.max_record_payload.min(DEFAULT_BULK_RECORD_PAYLOAD),
395        host_to_guest_credit_limit: DEFAULT_BULK_WINDOW,
396        guest_to_host_credit_limit: offer.guest_to_host_credit_limit,
397    })
398}
399
400/// Receive lifecycle data without making a non-bulk session's select loop spin.
401async fn recv_optional_mpsc<T>(receiver: &mut Option<mpsc::Receiver<T>>) -> Option<T> {
402    loop {
403        let Some(active) = receiver.as_mut() else {
404            return std::future::pending().await;
405        };
406        match active.recv().await {
407            Some(value) => return Some(value),
408            None => *receiver = None,
409        }
410    }
411}
412
413/// Return only the newest absolute credit update; intermediate grants are intentionally coalesced.
414async fn recv_optional_credit(
415    receiver: &mut Option<watch::Receiver<Option<BulkCredit>>>,
416) -> Option<BulkCredit> {
417    loop {
418        let Some(active) = receiver.as_mut() else {
419            return std::future::pending().await;
420        };
421        if active.changed().await.is_err() {
422            *receiver = None;
423            continue;
424        }
425        if let Some(credit) = *active.borrow_and_update() {
426            return Some(credit);
427        }
428    }
429}
430
431/// Apply an overtaking half-close only after the data worker has consumed every prior record.
432async fn apply_pending_tcp_finish<W>(
433    stream: &mut W,
434    state: &mut TcpBulkState,
435    pending: &mut Option<BulkFinish>,
436) -> Result<(), String>
437where
438    W: AsyncWrite + Unpin,
439{
440    let Some(finish) = *pending else {
441        return Ok(());
442    };
443    if finish.kind != BulkKind::Tcp || finish.flow != BulkFlow::HostToGuest {
444        return Err("bulk finish does not describe the TCP host-to-guest flow".into());
445    }
446    if finish.final_offset > state.receive.next_expected_offset() {
447        return Ok(());
448    }
449    state
450        .receive
451        .accept_finish(finish)
452        .map_err(|error| error.to_string())?;
453    stream
454        .shutdown()
455        .await
456        .map_err(|error| format!("shutdown TCP stream: {error}"))?;
457    *pending = None;
458    Ok(())
459}
460
461async fn relay_tcp_session(
462    id: u32,
463    stream: TcpStream,
464    mut commands: mpsc::Receiver<TcpCommand>,
465    bulk_control: Option<TcpBulkControlReceivers>,
466    tx: SessionOutputSender,
467    mut bulk: Option<TcpBulkState>,
468) {
469    // A single task still owns protocol state, but independent socket halves let a blocked
470    // destination write make progress concurrently with guest-to-host reads.
471    let (mut reader, mut writer) = stream.into_split();
472    let (mut credit_rx, mut finish_rx) = match bulk_control {
473        Some(control) => (Some(control.credit), Some(control.finish)),
474        None => (None, None),
475    };
476    let mut pending_finish = None;
477    let read_capacity = bulk.as_ref().map_or(TCP_CHUNK_SIZE, |state| {
478        state.send.max_record_payload() as usize
479    });
480    let mut read_buf = vec![0u8; read_capacity];
481    let mut terminal_sent = false;
482    let mut pending_write: Option<PendingTcpWrite> = None;
483    let mut write_shutdown = false;
484    // The destination half-closed its write side. We stop reading but keep the
485    // loop alive so host->destination data still flows until the host closes.
486    let mut read_eof = false;
487
488    loop {
489        // One EOF leaves the opposite half usable. Once both halves finish, all ordered
490        // writes have completed and the peer's final output/EOF is already queued. Exit so
491        // the existing terminal frame releases the host route and guest session together.
492        if read_eof && write_shutdown {
493            break;
494        }
495        let read_limit = bulk.as_ref().map_or(TCP_CHUNK_SIZE, |state| {
496            state
497                .send
498                .available_credit()
499                .min(state.send.max_record_payload() as u64) as usize
500        });
501        tokio::select! {
502            Some(finish) = recv_optional_mpsc(&mut finish_rx) => {
503                if pending_finish.replace(finish).is_some() {
504                    terminal_sent = send_tcp_failure(
505                        id,
506                        "duplicate TCP bulk finish".into(),
507                        &tx,
508                    )
509                    .await;
510                    break;
511                }
512                let Some(state) = bulk.as_mut() else {
513                    terminal_sent = send_tcp_failure(
514                        id,
515                        "bulk finish received on a generation-6 TCP stream".into(),
516                        &tx,
517                    )
518                    .await;
519                    break;
520                };
521                if pending_write.is_none() {
522                    if let Err(error) = apply_pending_tcp_finish(
523                        &mut writer,
524                        state,
525                        &mut pending_finish,
526                    ).await {
527                        terminal_sent = send_tcp_failure(
528                            id,
529                            format!("invalid TCP bulk finish: {error}"),
530                            &tx,
531                        )
532                        .await;
533                        break;
534                    }
535                    write_shutdown = pending_finish.is_none();
536                }
537            }
538            Some(credit) = recv_optional_credit(&mut credit_rx) => {
539                let Some(state) = bulk.as_mut() else {
540                    terminal_sent = send_tcp_failure(
541                        id,
542                        "bulk credit received on a generation-6 TCP stream".into(),
543                        &tx,
544                    )
545                    .await;
546                    break;
547                };
548                if let Err(error) = state.send.apply_credit(credit) {
549                    terminal_sent = send_tcp_failure(
550                        id,
551                        format!("invalid TCP bulk credit: {error}"),
552                        &tx,
553                    )
554                    .await;
555                    break;
556                }
557            }
558            read = reader.read(&mut read_buf[..read_limit]), if !read_eof && read_limit != 0 => {
559                match read {
560                    Ok(0) => {
561                        if let Some(state) = bulk.as_mut() {
562                            match state.send.finish() {
563                                Ok(finish) => {
564                                    send_raw_tcp_message(
565                                        id,
566                                        MessageType::BulkFinish,
567                                        &finish,
568                                        RawActivity::guest_message(),
569                                        None,
570                                        &tx,
571                                    )
572                                    .await;
573                                }
574                                Err(error) => {
575                                    eprintln!("failed to finish TCP bulk receive flow {id}: {error}");
576                                    break;
577                                }
578                            }
579                        } else {
580                            send_raw_tcp_message(
581                                id,
582                                MessageType::TcpEof,
583                                &TcpEof {},
584                                RawActivity::guest_message(),
585                                None,
586                                &tx,
587                            )
588                            .await;
589                        }
590                        read_eof = true;
591                    }
592                    Ok(n) => {
593                        if let Some(state) = bulk.as_mut() {
594                            let offset = match state.send.admit(n) {
595                                Ok(offset) => offset,
596                                Err(error) => {
597                                    eprintln!("failed to admit TCP bulk record {id}: {error}");
598                                    break;
599                                }
600                            };
601                            let Some(permit) = tx.reserve_bulk(n).await else {
602                                break;
603                            };
604                            let record = BulkRecord {
605                                id,
606                                kind: BulkKind::Tcp,
607                                flow: BulkFlow::GuestToHost,
608                                offset,
609                                payload: Bytes::copy_from_slice(&read_buf[..n]),
610                            };
611                            if !tx
612                                .send_reserved(
613                                    id,
614                                    SessionOutput::Bulk(BulkSessionOutput::new(
615                                        record,
616                                        RawActivity::tcp_bytes(n),
617                                    )),
618                                    permit,
619                                )
620                                .await
621                            {
622                                break;
623                            }
624                        } else {
625                            let Some(permit) = tx.reserve(TCP_OUTPUT_RESERVATION).await else {
626                                break;
627                            };
628                            let data = read_buf[..n].to_vec();
629                            if !send_raw_tcp_data(id, data, n, permit, &tx).await {
630                                break;
631                            }
632                        }
633                    }
634                    Err(e) => {
635                        terminal_sent = send_raw_tcp_message(
636                            id,
637                            MessageType::TcpFailed,
638                            &TcpFailed {
639                                error: format!("read TCP stream: {e}"),
640                            },
641                            RawActivity::guest_message(),
642                            Some(RawSessionCompletion::Tcp),
643                            &tx,
644                        )
645                        .await;
646                        break;
647                    }
648                }
649            }
650            write = async {
651                let pending = pending_write.as_ref().expect("guarded pending TCP write");
652                writer.write(&pending.payload[pending.written..]).await
653            }, if pending_write.is_some() => {
654                match write {
655                    Ok(0) => {
656                        terminal_sent = send_tcp_failure(
657                            id,
658                            "write TCP stream made no progress".into(),
659                            &tx,
660                        )
661                        .await;
662                        break;
663                    }
664                    Ok(written) => {
665                        let pending = pending_write.as_mut().expect("guarded pending TCP write");
666                        pending.written += written;
667                        if pending.written != pending.payload.len() {
668                            continue;
669                        }
670
671                        let completed = pending_write.take().expect("completed TCP write exists");
672                        let bulk_end = completed.bulk_end;
673                        // The destination socket has consumed the full payload. Release aggregate
674                        // input capacity before an outbound credit waits on the opposite lane.
675                        drop(completed._bulk_input_permit);
676                        drop(completed._control_input_charge);
677                        if let Some(end) = bulk_end {
678                            let Some(state) = bulk.as_mut() else {
679                                terminal_sent = send_tcp_failure(
680                                    id,
681                                    "bulk TCP write lost its protocol state".into(),
682                                    &tx,
683                                )
684                                .await;
685                                break;
686                            };
687                            match state.receive.consume(end) {
688                                Ok(Some(credit)) => {
689                                    if !send_raw_tcp_message(
690                                        id,
691                                        MessageType::BulkCredit,
692                                        &credit,
693                                        RawActivity::guest_message(),
694                                        None,
695                                        &tx,
696                                    )
697                                    .await
698                                    {
699                                        break;
700                                    }
701                                }
702                                Ok(None) => {}
703                                Err(error) => {
704                                    terminal_sent = send_tcp_failure(
705                                        id,
706                                        format!("advance TCP bulk credit: {error}"),
707                                        &tx,
708                                    )
709                                    .await;
710                                    break;
711                                }
712                            }
713                            let finish_was_pending = pending_finish.is_some();
714                            if let Err(error) = apply_pending_tcp_finish(
715                                &mut writer,
716                                state,
717                                &mut pending_finish,
718                            ).await {
719                                terminal_sent = send_tcp_failure(
720                                    id,
721                                    format!("invalid TCP bulk finish: {error}"),
722                                    &tx,
723                                )
724                                .await;
725                                break;
726                            }
727                            if finish_was_pending && pending_finish.is_none() {
728                                write_shutdown = true;
729                            }
730                        }
731                    }
732                    Err(error) => {
733                        terminal_sent = send_raw_tcp_message(
734                            id,
735                            MessageType::TcpFailed,
736                            &TcpFailed {
737                                error: format!("write TCP stream: {error}"),
738                            },
739                            RawActivity::guest_message(),
740                            Some(RawSessionCompletion::Tcp),
741                            &tx,
742                        )
743                        .await;
744                        break;
745                    }
746                }
747            }
748            command = commands.recv(), if pending_write.is_none() && !write_shutdown => {
749                match command {
750                    Some(TcpCommand::Data(data, charge)) => {
751                        if data.is_empty() {
752                            // An empty data message is not EOF and owns no socket write. Its
753                            // frame token still bounded admission until it reached this turn.
754                            drop(charge);
755                            continue;
756                        }
757                        if bulk.is_some() {
758                            terminal_sent = send_tcp_failure(
759                                id,
760                                "CBOR TCP data received after raw bulk acceptance".into(),
761                                &tx,
762                            )
763                            .await;
764                            break;
765                        }
766                        pending_write = Some(PendingTcpWrite {
767                            payload: Bytes::from(data),
768                            written: 0,
769                            bulk_end: None,
770                            _bulk_input_permit: None,
771                            _control_input_charge: charge,
772                        });
773                    }
774                    Some(TcpCommand::Eof(charge)) => {
775                        if bulk.is_some() {
776                            terminal_sent = send_tcp_failure(
777                                id,
778                                "CBOR TCP EOF received after raw bulk acceptance".into(),
779                                &tx,
780                            )
781                            .await;
782                            break;
783                        }
784                        if let Err(e) = writer.shutdown().await {
785                            terminal_sent = send_raw_tcp_message(
786                                id,
787                                MessageType::TcpFailed,
788                                &TcpFailed {
789                                    error: format!("shutdown TCP stream: {e}"),
790                                },
791                                RawActivity::guest_message(),
792                                Some(RawSessionCompletion::Tcp),
793                                &tx,
794                            )
795                            .await;
796                            break;
797                        }
798                        write_shutdown = true;
799                        drop(charge);
800                    }
801                    None => {
802                        break;
803                    }
804                    Some(TcpCommand::BulkRecord(record)) => {
805                        let Some(state) = bulk.as_mut() else {
806                            terminal_sent = send_tcp_failure(
807                                id,
808                                "raw bulk record received on a generation-6 TCP stream".into(),
809                                &tx,
810                            )
811                            .await;
812                            break;
813                        };
814                        let end = match state.receive.accept_record(record.record()) {
815                            Ok(end) => end,
816                            Err(error) => {
817                                terminal_sent = send_tcp_failure(
818                                    id,
819                                    format!("invalid TCP bulk record: {error}"),
820                                    &tx,
821                                )
822                                .await;
823                                break;
824                            }
825                        };
826                        let (record, permit) = record.into_parts();
827                        pending_write = Some(PendingTcpWrite {
828                            payload: record.payload,
829                            written: 0,
830                            bulk_end: Some(end),
831                            _bulk_input_permit: Some(permit),
832                            _control_input_charge: None,
833                        });
834                    }
835                }
836            }
837        }
838    }
839
840    if !terminal_sent {
841        send_raw_tcp_message(
842            id,
843            MessageType::TcpClosed,
844            &TcpClosed {},
845            RawActivity::guest_message(),
846            Some(RawSessionCompletion::Tcp),
847            &tx,
848        )
849        .await;
850    }
851}
852
853async fn send_tcp_failure(id: u32, error: String, tx: &SessionOutputSender) -> bool {
854    send_raw_tcp_message(
855        id,
856        MessageType::TcpFailed,
857        &TcpFailed { error },
858        RawActivity::guest_message(),
859        Some(RawSessionCompletion::Tcp),
860        tx,
861    )
862    .await
863}
864
865fn encode_tcp_message<T: serde::Serialize>(
866    id: u32,
867    t: MessageType,
868    payload: &T,
869    out_buf: &mut Vec<u8>,
870) -> Result<(), String> {
871    let msg = Message::with_payload(t, id, payload).map_err(|e| format!("encode tcp: {e}"))?;
872    codec::encode_to_buf(&msg, out_buf).map_err(|e| format!("encode tcp frame: {e}"))?;
873    Ok(())
874}
875
876async fn send_raw_tcp_message<T: serde::Serialize>(
877    id: u32,
878    t: MessageType,
879    payload: &T,
880    activity: RawActivity,
881    completion: Option<RawSessionCompletion>,
882    tx: &SessionOutputSender,
883) -> bool {
884    let mut buf = Vec::new();
885    match encode_tcp_message(id, t, payload, &mut buf) {
886        Ok(()) => {
887            tx.send(
888                id,
889                SessionOutput::Raw(RawSessionOutput::new(buf, activity, completion)),
890            )
891            .await
892        }
893        Err(e) => {
894            eprintln!("failed to encode tcp message for {id}: {e}");
895            false
896        }
897    }
898}
899
900/// Encode a TCP data event only after its retained allocation has reserved capacity.
901async fn send_raw_tcp_data(
902    id: u32,
903    data: Vec<u8>,
904    byte_count: usize,
905    permit: SessionOutputPermit,
906    tx: &SessionOutputSender,
907) -> bool {
908    let mut buf = Vec::new();
909    match encode_tcp_message(id, MessageType::TcpData, &TcpData { data }, &mut buf) {
910        Ok(()) => {
911            tx.send_reserved(
912                id,
913                SessionOutput::Raw(RawSessionOutput::new(
914                    buf,
915                    RawActivity::tcp_bytes(byte_count),
916                    None,
917                )),
918                permit,
919            )
920            .await
921        }
922        Err(error) => {
923            eprintln!("failed to encode TCP data for {id}: {error}");
924            false
925        }
926    }
927}
928
929//--------------------------------------------------------------------------------------------------
930// Tests
931//--------------------------------------------------------------------------------------------------
932
933#[cfg(test)]
934mod tests {
935    use std::time::Duration;
936
937    use microsandbox_protocol::message::FLAG_TERMINAL;
938    use tokio::net::TcpListener;
939
940    use super::*;
941
942    #[test]
943    fn admitted_transport_window_fits_each_tcp_input_queue() {
944        use microsandbox_protocol::core::{
945            WORKLOAD_TRANSPORT_BULK_FRAMES, WORKLOAD_TRANSPORT_CONTROL_FRAMES,
946        };
947
948        // Data and EOF retain their admission token until the socket consumes them. One input
949        // frame occupies at most one command slot, independent of its byte length.
950        assert!(
951            WORKLOAD_TRANSPORT_CONTROL_FRAMES + WORKLOAD_TRANSPORT_BULK_FRAMES
952                <= TCP_COMMAND_CAPACITY as u64
953        );
954    }
955
956    #[tokio::test]
957    async fn blocked_tcp_retains_data_and_eof_credit_until_consumption_or_cancel() {
958        use crate::serial::{InputLane, InputWindow};
959        use microsandbox_protocol::core::WorkloadTransportCredit;
960        use std::os::fd::AsRawFd;
961
962        for cancel in [false, true] {
963            let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
964            // Bound receive buffering before accept so the peer cannot consume the entire
965            // admitted8MiB while the application intentionally has not started reading.
966            let receive_bytes: libc::c_int = 64 * 1024;
967            assert_eq!(
968                unsafe {
969                    libc::setsockopt(
970                        listener.as_raw_fd(),
971                        libc::SOL_SOCKET,
972                        libc::SO_RCVBUF,
973                        (&receive_bytes as *const libc::c_int).cast(),
974                        std::mem::size_of_val(&receive_bytes) as libc::socklen_t,
975                    )
976                },
977                0
978            );
979            let (sender, mut output) = SessionOutputSender::channel();
980            let session = TcpSession::open(
981                8,
982                TcpConnect {
983                    host: "127.0.0.1".into(),
984                    port: listener.local_addr().unwrap().port(),
985                    bulk: None,
986                },
987                &sender,
988            );
989            let (mut peer, _) = listener.accept().await.unwrap();
990            assert_eq!(recv_message(&mut output).await.t, MessageType::TcpConnected);
991            let initial = WorkloadTransportCredit {
992                control_bytes: 64,
993                control_frames: 2,
994                bulk_bytes: 8 * 1024 * 1024,
995                bulk_frames: 2,
996            };
997            let ledger = InputWindow::new(initial);
998            let payload_len = initial.bulk_bytes as usize - 64;
999            let data_charge = ledger.admit(InputLane::Bulk, payload_len + 32).unwrap();
1000            let eof_charge = ledger.admit(InputLane::Bulk, 32).unwrap();
1001            tokio::time::timeout(Duration::from_millis(100), async {
1002                session
1003                    .write_data_charged(vec![0x5c; payload_len], Some(data_charge))
1004                    .await
1005                    .unwrap();
1006                session.close_write_charged(Some(eof_charge)).await.unwrap();
1007            })
1008            .await
1009            .expect("admitted input waited for a blocked TCP consumer");
1010            tokio::time::sleep(Duration::from_millis(20)).await;
1011            assert_eq!(ledger.credit().unwrap(), initial);
1012            assert!(ledger.admit(InputLane::Bulk, 1).is_err());
1013            if cancel {
1014                session.close();
1015                wait_finished(&session).await;
1016            } else {
1017                let mut bytes = Vec::new();
1018                tokio::time::timeout(Duration::from_secs(10), peer.read_to_end(&mut bytes))
1019                    .await
1020                    .expect("ordered TCP EOF did not arrive")
1021                    .unwrap();
1022                assert_eq!(bytes.len(), payload_len);
1023                assert!(bytes.iter().all(|byte| *byte == 0x5c));
1024                session.close();
1025                wait_finished(&session).await;
1026            }
1027            assert_eq!(ledger.credit().unwrap().bulk_bytes, initial.bulk_bytes * 2);
1028            assert_eq!(ledger.credit().unwrap().bulk_frames, 4);
1029            assert_eq!(
1030                ledger.credit().unwrap().control_bytes,
1031                initial.control_bytes
1032            );
1033        }
1034    }
1035
1036    #[tokio::test]
1037    async fn connect_failure_sends_terminal_failed() {
1038        let (session_tx, mut session_rx) = SessionOutputSender::channel();
1039
1040        let session = TcpSession::open(
1041            7,
1042            TcpConnect {
1043                host: "127.0.0.1".to_string(),
1044                port: 0,
1045                bulk: None,
1046            },
1047            &session_tx,
1048        );
1049
1050        // The connect runs in the task and reports failure over session_tx.
1051        let msg = recv_message(&mut session_rx).await;
1052        assert_eq!(msg.t, MessageType::TcpFailed);
1053        assert_eq!(msg.flags, FLAG_TERMINAL);
1054        let failed: TcpFailed = msg.payload().unwrap();
1055        assert!(failed.error.contains("connect 127.0.0.1:0"));
1056
1057        wait_finished(&session).await;
1058    }
1059
1060    #[tokio::test]
1061    async fn close_request_finishes_session_task() {
1062        let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1063        let port = listener.local_addr().unwrap().port();
1064        let (session_tx, mut session_rx) = SessionOutputSender::channel();
1065        let accept_task = tokio::spawn(async move {
1066            let (_socket, _) = listener.accept().await.unwrap();
1067            tokio::time::sleep(Duration::from_secs(5)).await;
1068        });
1069
1070        let session = TcpSession::open(
1071            9,
1072            TcpConnect {
1073                host: "127.0.0.1".to_string(),
1074                port,
1075                bulk: None,
1076            },
1077            &session_tx,
1078        );
1079
1080        let connected = recv_message(&mut session_rx).await;
1081        assert_eq!(connected.t, MessageType::TcpConnected);
1082
1083        session.close();
1084        wait_finished(&session).await;
1085
1086        accept_task.abort();
1087    }
1088
1089    #[tokio::test]
1090    async fn destination_eof_keeps_session_open_for_host_writes() {
1091        let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1092        let port = listener.local_addr().unwrap().port();
1093        let (session_tx, mut session_rx) = SessionOutputSender::channel();
1094
1095        // The destination half-closes its write side, then keeps reading so it
1096        // still receives whatever the host sends after the EOF.
1097        let (got_tx, got_rx) = tokio::sync::oneshot::channel();
1098        let accept_task = tokio::spawn(async move {
1099            let (mut socket, _) = listener.accept().await.unwrap();
1100            socket.shutdown().await.unwrap();
1101            let mut buf = Vec::new();
1102            socket.read_to_end(&mut buf).await.unwrap();
1103            let _ = got_tx.send(buf);
1104        });
1105
1106        let session = TcpSession::open(
1107            11,
1108            TcpConnect {
1109                host: "127.0.0.1".to_string(),
1110                port,
1111                bulk: None,
1112            },
1113            &session_tx,
1114        );
1115
1116        let connected = recv_message(&mut session_rx).await;
1117        assert_eq!(connected.t, MessageType::TcpConnected);
1118
1119        // The destination's FIN surfaces as a non-terminal TcpEof, and the
1120        // session stays alive.
1121        let eof = recv_message(&mut session_rx).await;
1122        assert_eq!(eof.t, MessageType::TcpEof);
1123        assert_ne!(eof.flags, FLAG_TERMINAL);
1124        assert!(!session.is_finished());
1125
1126        // The host can still reach the destination after that EOF.
1127        session.write_data(b"after-eof".to_vec()).await.unwrap();
1128        session.close_write().await.unwrap();
1129        let received = tokio::time::timeout(Duration::from_secs(1), got_rx)
1130            .await
1131            .unwrap()
1132            .unwrap();
1133        assert_eq!(received, b"after-eof");
1134
1135        // An explicit close tears the session down.
1136        session.close();
1137        wait_finished(&session).await;
1138
1139        accept_task.await.unwrap();
1140    }
1141
1142    #[tokio::test]
1143    async fn active_raw_credit_validation_and_inline_negotiation_still_apply() {
1144        for raw in [false, true] {
1145            let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1146            let (tx, mut rx) = SessionOutputSender::channel();
1147            let session = TcpSession::open(
1148                41,
1149                TcpConnect {
1150                    host: "127.0.0.1".into(),
1151                    port: listener.local_addr().unwrap().port(),
1152                    bulk: raw.then(BulkOffer::tcp),
1153                },
1154                &tx,
1155            );
1156            let (_peer, _) = listener.accept().await.unwrap();
1157            assert_eq!(recv_message(&mut rx).await.t, MessageType::TcpConnected);
1158            if raw {
1159                assert_eq!(recv_message(&mut rx).await.t, MessageType::BulkAccepted);
1160            }
1161            let result = session
1162                .apply_credit(BulkCredit {
1163                    kind: BulkKind::Tcp,
1164                    flow: BulkFlow::GuestToHost,
1165                    consumed_offset: 1,
1166                    credit_limit: DEFAULT_BULK_WINDOW + 1,
1167                })
1168                .await;
1169            if raw {
1170                result.unwrap();
1171                let failed = tokio::time::timeout(Duration::from_secs(1), recv_message(&mut rx))
1172                    .await
1173                    .unwrap();
1174                assert_eq!(failed.t, MessageType::TcpFailed);
1175                assert_eq!(failed.flags, FLAG_TERMINAL);
1176                assert!(
1177                    failed
1178                        .payload::<TcpFailed>()
1179                        .unwrap()
1180                        .error
1181                        .contains("not admitted")
1182                );
1183                wait_finished(&session).await;
1184            } else {
1185                assert!(result.unwrap_err().contains("generation-6"));
1186                session.close();
1187                wait_finished(&session).await;
1188            }
1189        }
1190    }
1191
1192    #[tokio::test]
1193    async fn both_half_close_orders_preserve_data_and_emit_one_terminal() {
1194        for raw in [false, true] {
1195            for peer_first in [false, true] {
1196                tokio::time::timeout(Duration::from_secs(5), async {
1197                    let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1198                    let (tx, mut rx) = SessionOutputSender::channel();
1199                    let session = TcpSession::open(
1200                        31,
1201                        TcpConnect {
1202                            host: "127.0.0.1".into(),
1203                            port: listener.local_addr().unwrap().port(),
1204                            bulk: raw.then(BulkOffer::tcp),
1205                        },
1206                        &tx,
1207                    );
1208                    let (mut peer, _) = listener.accept().await.unwrap();
1209                    assert_eq!(recv_message(&mut rx).await.t, MessageType::TcpConnected);
1210                    if raw {
1211                        assert_eq!(recv_message(&mut rx).await.t, MessageType::BulkAccepted);
1212                    }
1213                    let host_data = b"host data survives the peer's first EOF";
1214                    let peer_data = b"peer data survives the host's first EOF";
1215                    if peer_first {
1216                        peer.write_all(peer_data).await.unwrap();
1217                        peer.shutdown().await.unwrap();
1218                        assert_tcp_output_through_eof(&mut rx, raw, peer_data).await;
1219                        assert!(!session.is_finished(), "one EOF must preserve host writes");
1220                        send_test_input_and_eof(&session, raw, host_data).await;
1221                    } else {
1222                        send_test_input_and_eof(&session, raw, host_data).await;
1223                    }
1224
1225                    let mut received = Vec::new();
1226                    peer.read_to_end(&mut received).await.unwrap();
1227                    assert_eq!(received, host_data);
1228                    if !peer_first {
1229                        assert!(!session.is_finished(), "one EOF must preserve peer output");
1230                        peer.write_all(peer_data).await.unwrap();
1231                        peer.shutdown().await.unwrap();
1232                        assert_tcp_output_through_eof(&mut rx, raw, peer_data).await;
1233                    }
1234                    assert_one_normal_terminal(&session, &mut rx).await;
1235                })
1236                .await
1237                .unwrap_or_else(|_| {
1238                    panic!("TCP completion timed out: raw={raw}, peer_first={peer_first}")
1239                });
1240            }
1241        }
1242    }
1243
1244    #[tokio::test]
1245    async fn raw_finish_waits_for_delayed_record_and_pending_socket_write_before_terminal() {
1246        use std::os::fd::AsRawFd;
1247
1248        tokio::time::timeout(Duration::from_secs(5), async {
1249            let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1250            let (stream, accepted) = tokio::join!(
1251                TcpStream::connect(listener.local_addr().unwrap()),
1252                listener.accept(),
1253            );
1254            let stream = stream.unwrap();
1255            let (mut peer, _) = accepted.unwrap();
1256            // Make the last record larger than both fixed socket buffers. The test observes a
1257            // delivered prefix before draining the rest, so EOF cannot be credited at enqueue.
1258            for (fd, option, bytes) in [
1259                (stream.as_raw_fd(), libc::SO_SNDBUF, 4096 as libc::c_int),
1260                (peer.as_raw_fd(), libc::SO_RCVBUF, 65536 as libc::c_int),
1261            ] {
1262                assert_eq!(
1263                    unsafe {
1264                        libc::setsockopt(
1265                            fd,
1266                            libc::SOL_SOCKET,
1267                            option,
1268                            (&bytes as *const libc::c_int).cast(),
1269                            std::mem::size_of_val(&bytes) as libc::socklen_t,
1270                        )
1271                    },
1272                    0
1273                );
1274            }
1275            let (tx, mut rx) = SessionOutputSender::channel();
1276            let (commands, commands_rx) = mpsc::channel(TCP_COMMAND_CAPACITY);
1277            let (credit, credit_rx) = watch::channel(None);
1278            let (finish, finish_rx) = mpsc::channel(1);
1279            let task = tokio::spawn(relay_tcp_session(
1280                37,
1281                stream,
1282                commands_rx,
1283                Some(TcpBulkControlReceivers {
1284                    credit: credit_rx,
1285                    finish: finish_rx,
1286                }),
1287                tx,
1288                Some(TcpBulkState {
1289                    send: BulkSendState::new(
1290                        BulkKind::Tcp,
1291                        BulkFlow::GuestToHost,
1292                        DEFAULT_BULK_RECORD_PAYLOAD,
1293                        DEFAULT_BULK_WINDOW,
1294                    )
1295                    .unwrap(),
1296                    receive: BulkReceiveState::new(
1297                        BulkKind::Tcp,
1298                        BulkFlow::HostToGuest,
1299                        DEFAULT_BULK_RECORD_PAYLOAD,
1300                        DEFAULT_BULK_WINDOW,
1301                        DEFAULT_BULK_WINDOW,
1302                    )
1303                    .unwrap(),
1304                }),
1305            ));
1306            let session = TcpSession {
1307                owner_id: 37,
1308                commands,
1309                bulk_control: Some(TcpBulkControlSenders { credit, finish }),
1310                task,
1311                bulk: true,
1312            };
1313            peer.shutdown().await.unwrap();
1314            assert_tcp_output_through_eof(&mut rx, true, b"").await;
1315            let payload = Bytes::from(vec![0x6a; DEFAULT_BULK_RECORD_PAYLOAD as usize]);
1316            session
1317                .finish_bulk(BulkFinish {
1318                    kind: BulkKind::Tcp,
1319                    flow: BulkFlow::HostToGuest,
1320                    final_offset: payload.len() as u64,
1321                })
1322                .await
1323                .unwrap();
1324            while session.bulk_control.as_ref().unwrap().finish.capacity() == 0 {
1325                tokio::task::yield_now().await;
1326            }
1327            assert!(
1328                !session.is_finished(),
1329                "finish cannot skip its missing final record"
1330            );
1331            assert!(matches!(
1332                rx.try_recv(),
1333                Err(mpsc::error::TryRecvError::Empty)
1334            ));
1335            session
1336                .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1337                    id: 37,
1338                    kind: BulkKind::Tcp,
1339                    flow: BulkFlow::HostToGuest,
1340                    offset: 0,
1341                    payload: payload.clone(),
1342                }))
1343                .await
1344                .unwrap();
1345            let mut received = vec![0];
1346            peer.read_exact(&mut received).await.unwrap();
1347            assert!(
1348                !session.is_finished(),
1349                "finish cannot skip a partial socket write"
1350            );
1351            assert!(matches!(
1352                rx.try_recv(),
1353                Err(mpsc::error::TryRecvError::Empty)
1354            ));
1355            peer.read_to_end(&mut received).await.unwrap();
1356            assert_eq!(received, payload);
1357            assert_one_normal_terminal(&session, &mut rx).await;
1358        })
1359        .await
1360        .expect("delayed raw record did not finish normally");
1361    }
1362
1363    #[tokio::test]
1364    async fn raw_bulk_tcp_relays_both_directions_and_exact_half_closes() {
1365        let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1366        let port = listener.local_addr().unwrap().port();
1367        let (session_tx, mut session_rx) = SessionOutputSender::channel();
1368        let (got_tx, got_rx) = tokio::sync::oneshot::channel();
1369        let accept_task = tokio::spawn(async move {
1370            let (mut socket, _) = listener.accept().await.unwrap();
1371            socket.write_all(b"from-destination").await.unwrap();
1372            socket.shutdown().await.unwrap();
1373            let mut received = Vec::new();
1374            socket.read_to_end(&mut received).await.unwrap();
1375            got_tx.send(received).unwrap();
1376        });
1377
1378        let session = TcpSession::open(
1379            13,
1380            TcpConnect {
1381                host: "127.0.0.1".to_string(),
1382                port,
1383                bulk: Some(BulkOffer::tcp()),
1384            },
1385            &session_tx,
1386        );
1387        assert_eq!(
1388            recv_message(&mut session_rx).await.t,
1389            MessageType::TcpConnected
1390        );
1391        assert_eq!(
1392            recv_message(&mut session_rx).await.t,
1393            MessageType::BulkAccepted
1394        );
1395
1396        let host_payload = Bytes::from_static(b"from-host");
1397        session
1398            .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1399                id: 13,
1400                kind: BulkKind::Tcp,
1401                flow: BulkFlow::HostToGuest,
1402                offset: 0,
1403                payload: host_payload.clone(),
1404            }))
1405            .await
1406            .unwrap();
1407        session
1408            .finish_bulk(BulkFinish {
1409                kind: BulkKind::Tcp,
1410                flow: BulkFlow::HostToGuest,
1411                final_offset: host_payload.len() as u64,
1412            })
1413            .await
1414            .unwrap();
1415
1416        let record = recv_bulk(&mut session_rx).await;
1417        assert_eq!(record.flow, BulkFlow::GuestToHost);
1418        assert_eq!(record.offset, 0);
1419        assert_eq!(record.payload, Bytes::from_static(b"from-destination"));
1420        let finish = recv_message(&mut session_rx).await;
1421        assert_eq!(finish.t, MessageType::BulkFinish);
1422        let finish: BulkFinish = finish.payload().unwrap();
1423        assert_eq!(finish.final_offset, b"from-destination".len() as u64);
1424
1425        let received = tokio::time::timeout(Duration::from_secs(1), got_rx)
1426            .await
1427            .unwrap()
1428            .unwrap();
1429        assert_eq!(received, host_payload);
1430        session.close();
1431        wait_finished(&session).await;
1432        accept_task.await.unwrap();
1433    }
1434
1435    #[tokio::test]
1436    async fn blocked_host_to_guest_write_does_not_stop_guest_to_host_reads() {
1437        let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1438        let port = listener.local_addr().unwrap().port();
1439        let (session_tx, mut session_rx) = SessionOutputSender::channel();
1440        let host_payload = Bytes::from(vec![0x3c; DEFAULT_BULK_RECORD_PAYLOAD as usize]);
1441        let guest_payload = vec![0x7a; DEFAULT_BULK_WINDOW as usize];
1442        let expected_guest_payload = guest_payload.clone();
1443        let expected_host_payload = host_payload.clone();
1444        let host_payload_len = host_payload.len();
1445        let (got_tx, got_rx) = tokio::sync::oneshot::channel();
1446
1447        // Both peers deliberately fill their send side before reading the other direction. A
1448        // single write_all/read loop deadlocks here once the kernel buffers fill; split halves do
1449        // not, because agentd keeps draining the destination while its write is pending.
1450        let accept_task = tokio::spawn(async move {
1451            let (mut socket, _) = listener.accept().await.unwrap();
1452            let mut received = vec![0u8; host_payload_len];
1453            socket.read_exact(&mut received[..1]).await.unwrap();
1454            socket.write_all(&guest_payload).await.unwrap();
1455            socket.read_exact(&mut received[1..]).await.unwrap();
1456            got_tx.send(received).unwrap();
1457        });
1458
1459        let session = TcpSession::open(
1460            17,
1461            TcpConnect {
1462                host: "127.0.0.1".to_string(),
1463                port,
1464                bulk: Some(BulkOffer::tcp()),
1465            },
1466            &session_tx,
1467        );
1468        assert_eq!(
1469            recv_message(&mut session_rx).await.t,
1470            MessageType::TcpConnected
1471        );
1472        assert_eq!(
1473            recv_message(&mut session_rx).await.t,
1474            MessageType::BulkAccepted
1475        );
1476
1477        session
1478            .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1479                id: 17,
1480                kind: BulkKind::Tcp,
1481                flow: BulkFlow::HostToGuest,
1482                offset: 0,
1483                payload: host_payload,
1484            }))
1485            .await
1486            .unwrap();
1487
1488        let received_guest_payload = tokio::time::timeout(Duration::from_secs(5), async {
1489            let mut received = Vec::with_capacity(expected_guest_payload.len());
1490            while received.len() < expected_guest_payload.len() {
1491                let record = recv_bulk(&mut session_rx).await;
1492                assert_eq!(record.offset, received.len() as u64);
1493                received.extend_from_slice(&record.payload);
1494            }
1495            received
1496        })
1497        .await
1498        .expect("guest-to-host reads must progress while the opposite write is blocked");
1499        assert_eq!(received_guest_payload, expected_guest_payload);
1500
1501        let received_host_payload = tokio::time::timeout(Duration::from_secs(5), got_rx)
1502            .await
1503            .unwrap()
1504            .unwrap();
1505        assert_eq!(received_host_payload, expected_host_payload);
1506
1507        session.close();
1508        wait_finished(&session).await;
1509        accept_task.await.unwrap();
1510    }
1511
1512    #[tokio::test]
1513    async fn full_bulk_data_queue_does_not_starve_credit_or_finish() {
1514        let (commands, _commands_rx) = mpsc::channel(TCP_COMMAND_CAPACITY);
1515        let (credit, mut credit_rx) = watch::channel(None);
1516        let (finish, mut finish_rx) = mpsc::channel(1);
1517        let task = tokio::spawn(std::future::pending());
1518        let session = TcpSession {
1519            owner_id: 29,
1520            commands,
1521            bulk_control: Some(TcpBulkControlSenders { credit, finish }),
1522            task,
1523            bulk: true,
1524        };
1525        let payload = Bytes::from(vec![0u8; MIN_BULK_RECORD_PAYLOAD as usize]);
1526
1527        for index in 0..TCP_COMMAND_CAPACITY {
1528            session
1529                .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1530                    id: 29,
1531                    kind: BulkKind::Tcp,
1532                    flow: BulkFlow::HostToGuest,
1533                    offset: (index * MIN_BULK_RECORD_PAYLOAD as usize) as u64,
1534                    payload: payload.clone(),
1535                }))
1536                .await
1537                .unwrap();
1538        }
1539        assert_eq!(session.commands.capacity(), 0);
1540
1541        let credit_update = BulkCredit {
1542            kind: BulkKind::Tcp,
1543            flow: BulkFlow::GuestToHost,
1544            consumed_offset: 4,
1545            credit_limit: 8,
1546        };
1547        session.apply_credit(credit_update).await.unwrap();
1548        credit_rx.changed().await.unwrap();
1549        assert_eq!(*credit_rx.borrow_and_update(), Some(credit_update));
1550
1551        let finish_update = BulkFinish {
1552            kind: BulkKind::Tcp,
1553            flow: BulkFlow::HostToGuest,
1554            final_offset: DEFAULT_BULK_WINDOW,
1555        };
1556        session.finish_bulk(finish_update).await.unwrap();
1557        assert_eq!(finish_rx.recv().await, Some(finish_update));
1558
1559        session.close();
1560    }
1561
1562    async fn send_test_input_and_eof(session: &TcpSession, raw: bool, data: &[u8]) {
1563        if raw {
1564            session
1565                .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1566                    id: session.owner_id(),
1567                    kind: BulkKind::Tcp,
1568                    flow: BulkFlow::HostToGuest,
1569                    offset: 0,
1570                    payload: Bytes::copy_from_slice(data),
1571                }))
1572                .await
1573                .unwrap();
1574            session
1575                .finish_bulk(BulkFinish {
1576                    kind: BulkKind::Tcp,
1577                    flow: BulkFlow::HostToGuest,
1578                    final_offset: data.len() as u64,
1579                })
1580                .await
1581                .unwrap();
1582        } else {
1583            session.write_data(data.to_vec()).await.unwrap();
1584            session.close_write().await.unwrap();
1585        }
1586    }
1587
1588    async fn assert_tcp_output_through_eof(
1589        rx: &mut mpsc::Receiver<SessionOutputEnvelope>,
1590        raw: bool,
1591        expected: &[u8],
1592    ) {
1593        let mut received = Vec::new();
1594        loop {
1595            let envelope = rx.recv().await.expect("TCP output ended before EOF");
1596            match envelope.output {
1597                SessionOutput::Bulk(output) => {
1598                    assert!(raw);
1599                    assert_eq!(output.record.offset, received.len() as u64);
1600                    received.extend_from_slice(&output.record.payload);
1601                }
1602                SessionOutput::Raw(mut output) => {
1603                    let message = decode_one_message(&mut output.frame);
1604                    assert_eq!(message.flags & FLAG_TERMINAL, 0, "terminal preceded EOF");
1605                    match message.t {
1606                        MessageType::TcpData => {
1607                            assert!(!raw);
1608                            received.extend(message.payload::<TcpData>().unwrap().data);
1609                        }
1610                        MessageType::TcpEof => {
1611                            assert!(!raw);
1612                            break;
1613                        }
1614                        MessageType::BulkFinish => {
1615                            assert!(raw);
1616                            let finish = message.payload::<BulkFinish>().unwrap();
1617                            assert_eq!(finish.final_offset, received.len() as u64);
1618                            break;
1619                        }
1620                        _ => panic!("unexpected TCP output: {:?}", message.t),
1621                    }
1622                }
1623                _ => panic!("unexpected non-TCP output"),
1624            }
1625        }
1626        assert_eq!(received, expected);
1627    }
1628
1629    async fn assert_one_normal_terminal(
1630        session: &TcpSession,
1631        rx: &mut mpsc::Receiver<SessionOutputEnvelope>,
1632    ) {
1633        let closed = recv_message(rx).await;
1634        assert_eq!(closed.t, MessageType::TcpClosed);
1635        assert_eq!(closed.flags, FLAG_TERMINAL);
1636        closed.payload::<TcpClosed>().unwrap();
1637        wait_finished(session).await;
1638        assert!(matches!(
1639            rx.try_recv(),
1640            Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected)
1641        ));
1642    }
1643
1644    async fn wait_finished(session: &TcpSession) {
1645        tokio::time::timeout(Duration::from_secs(1), async {
1646            while !session.is_finished() {
1647                tokio::time::sleep(Duration::from_millis(10)).await;
1648            }
1649        })
1650        .await
1651        .unwrap();
1652    }
1653
1654    fn decode_one_message(buf: &mut Vec<u8>) -> Message {
1655        codec::try_decode_from_buf(buf).unwrap().unwrap()
1656    }
1657
1658    async fn recv_message(rx: &mut mpsc::Receiver<SessionOutputEnvelope>) -> Message {
1659        let envelope = rx.recv().await.unwrap();
1660        let SessionOutput::Raw(mut output) = envelope.output else {
1661            panic!("expected SessionOutput::Raw frame");
1662        };
1663        decode_one_message(&mut output.frame)
1664    }
1665
1666    async fn recv_bulk(rx: &mut mpsc::Receiver<SessionOutputEnvelope>) -> BulkRecord {
1667        let envelope = rx.recv().await.unwrap();
1668        let SessionOutput::Bulk(output) = envelope.output else {
1669            panic!("expected SessionOutput::Bulk record");
1670        };
1671        output.record
1672    }
1673}