Skip to main content

microsandbox_network/engine/tcp/
connection.rs

1//! Connection tracker: manages smoltcp TCP sockets for the poll loop.
2//!
3//! Creates sockets on SYN detection, tracks connection lifecycle, relays data
4//! between smoltcp sockets and proxy task channels, and cleans up closed
5//! connections.
6
7use std::collections::{HashMap, HashSet};
8use std::net::SocketAddr;
9use std::num::NonZeroUsize;
10use std::sync::Arc;
11use std::sync::atomic::{AtomicU8, Ordering};
12
13use bytes::Bytes;
14use smoltcp::iface::{SocketHandle, SocketSet};
15use smoltcp::socket::tcp;
16use smoltcp::wire::IpListenEndpoint;
17use tokio::sync::mpsc;
18
19use crate::tcp::deferred_close::DeferredClose;
20
21//--------------------------------------------------------------------------------------------------
22// Constants
23//--------------------------------------------------------------------------------------------------
24
25/// Log target for opt-in profiling events.
26const PROFILING_TARGET: &str = "microsandbox::profiling";
27
28/// TCP socket receive buffer size (64 KiB).
29const TCP_RX_BUF_SIZE: usize = 65536;
30
31/// TCP socket transmit buffer size (64 KiB).
32const TCP_TX_BUF_SIZE: usize = 65536;
33
34/// Capacity of the mpsc channels between the poll loop and proxy tasks.
35const CHANNEL_CAPACITY: usize = 32;
36
37/// Buffer size for reading from smoltcp sockets.
38const RELAY_BUF_SIZE: usize = 16384;
39
40//--------------------------------------------------------------------------------------------------
41// Types
42//--------------------------------------------------------------------------------------------------
43
44/// Terminal connection status reported by an outbound proxy task.
45#[repr(u8)]
46#[derive(Clone, Copy, Debug, Eq, PartialEq)]
47pub enum ProxyConnectStatus {
48    /// No final proxy connection status has been reported yet.
49    Pending = 0,
50    /// The proxy connected to the upstream.
51    Connected = 1,
52    /// The proxy denied the connection before dialing upstream.
53    PolicyDenied = 2,
54    /// The proxy attempted to dial upstream and the connect failed.
55    UpstreamConnectFailed = 3,
56}
57
58/// Shared status for an outbound proxy task.
59///
60/// The smoltcp poll loop reads this when the proxy task exits to decide
61/// whether the guest should see a clean close or a TCP reset.
62pub struct ProxyConnectState {
63    status: AtomicU8,
64}
65
66/// Tracks TCP connections between guest and proxy tasks.
67///
68/// Each guest TCP connection maps to a smoltcp socket and a pair of channels
69/// connecting it to a tokio proxy task. The tracker handles:
70///
71/// - **Socket creation** — on SYN detection, before smoltcp processes the frame.
72/// - **Data relay** — shuttles bytes between smoltcp sockets and channels.
73/// - **Lifecycle detection** — identifies newly-established connections for
74///   proxy spawning.
75/// - **Cleanup** — removes closed sockets from the socket set.
76pub struct TcpConnectionTracker {
77    /// Active connections keyed by smoltcp socket handle.
78    connections: HashMap<SocketHandle, Connection>,
79    /// Secondary index for O(1) duplicate-SYN detection by (src, dst) 4-tuple.
80    connection_keys: HashSet<(SocketAddr, SocketAddr)>,
81    /// Max concurrent connections (from NetworkConfig).
82    max_tcp_connections: Option<NonZeroUsize>,
83    rejected_connections: u64,
84}
85
86/// Deprecated name for [`TcpConnectionTracker`].
87#[deprecated(note = "use TcpConnectionTracker instead")]
88pub type ConnectionTracker = TcpConnectionTracker;
89
90/// Internal state for a single tracked TCP connection.
91struct Connection {
92    /// Guest source address (from the guest's SYN).
93    src: SocketAddr,
94    /// Original destination (from the guest's SYN).
95    dst: SocketAddr,
96    /// Sends data from smoltcp socket to proxy task (guest → server).
97    ///
98    /// Set to `None` once the guest half-closes (FIN) and all its data has
99    /// been relayed: dropping the sender makes the proxy task's
100    /// `from_smoltcp.recv()` return `None`, propagating the half-close
101    /// upstream while the server → guest direction stays open.
102    to_proxy: Option<mpsc::Sender<Bytes>>,
103    /// Receives data from proxy task to write to smoltcp socket (server → guest).
104    from_proxy: mpsc::Receiver<Bytes>,
105    /// Proxy-side channel ends, held until the connection is ESTABLISHED.
106    /// Taken by [`TcpConnectionTracker::take_new_connections()`].
107    proxy_channels: Option<ProxyChannels>,
108    /// Whether a proxy task has been spawned for this connection.
109    proxy_spawned: bool,
110    /// Status reported by the proxy task before it exits.
111    proxy_connect: Arc<ProxyConnectState>,
112    /// Partial data from proxy that couldn't be fully written to smoltcp socket.
113    write_buf: Option<(Bytes, usize)>,
114    /// Data read from smoltcp socket that couldn't be sent to proxy (channel full).
115    /// Must be sent before reading more from the socket to preserve stream order.
116    read_buf: Option<Bytes>,
117    /// Progress deadline while draining after the host task exits.
118    deferred_close: DeferredClose,
119    /// Egress policy already denied this flow at SYN time; the connection
120    /// was accepted only so an HTTP/HTTPS client can be answered with 403.
121    policy_denied: bool,
122}
123
124/// Proxy-side channel ends, created at socket creation time and taken when
125/// the connection becomes ESTABLISHED.
126struct ProxyChannels {
127    /// Receive data from smoltcp socket (guest → proxy task).
128    from_smoltcp: mpsc::Receiver<Bytes>,
129    /// Send data to smoltcp socket (proxy task → guest).
130    to_smoltcp: mpsc::Sender<Bytes>,
131}
132
133/// Information for spawning a proxy task for a newly established connection.
134///
135/// Returned by [`TcpConnectionTracker::take_new_connections()`]. The poll loop
136/// passes this to the proxy task spawner.
137pub struct NewConnection {
138    /// Original destination the guest was connecting to.
139    pub dst: SocketAddr,
140    /// Receive data from smoltcp socket (guest → proxy task).
141    pub from_smoltcp: mpsc::Receiver<Bytes>,
142    /// Send data to smoltcp socket (proxy task → guest).
143    pub to_smoltcp: mpsc::Sender<Bytes>,
144    /// Status the proxy task updates before it exits.
145    pub proxy_connect: Arc<ProxyConnectState>,
146    /// Egress policy already denied this flow at SYN time. The dispatcher
147    /// must answer it (HTTP 403) and never dial upstream.
148    pub policy_denied: bool,
149}
150
151//--------------------------------------------------------------------------------------------------
152// Methods
153//--------------------------------------------------------------------------------------------------
154
155impl ProxyConnectStatus {
156    fn as_u8(self) -> u8 {
157        self as u8
158    }
159
160    fn from_u8(value: u8) -> Self {
161        match value {
162            value if value == Self::Connected as u8 => Self::Connected,
163            value if value == Self::PolicyDenied as u8 => Self::PolicyDenied,
164            value if value == Self::UpstreamConnectFailed as u8 => Self::UpstreamConnectFailed,
165            _ => Self::Pending,
166        }
167    }
168}
169
170impl ProxyConnectState {
171    /// Create a new pending proxy connection status.
172    pub fn new() -> Self {
173        Self {
174            status: AtomicU8::new(ProxyConnectStatus::Pending.as_u8()),
175        }
176    }
177
178    /// Mark the proxy as successfully connected to upstream.
179    pub fn mark_connected(&self) {
180        self.store(ProxyConnectStatus::Connected);
181    }
182
183    /// Mark the proxy as denied by egress policy before dialing upstream.
184    pub fn mark_policy_denied(&self) {
185        self.store(ProxyConnectStatus::PolicyDenied);
186    }
187
188    /// Mark the proxy as failed while dialing upstream.
189    pub fn mark_upstream_connect_failed(&self) {
190        self.store(ProxyConnectStatus::UpstreamConnectFailed);
191    }
192
193    /// Load the latest proxy connection status.
194    pub fn status(&self) -> ProxyConnectStatus {
195        ProxyConnectStatus::from_u8(self.status.load(Ordering::Acquire))
196    }
197
198    fn store(&self, status: ProxyConnectStatus) {
199        self.status.store(status.as_u8(), Ordering::Release);
200    }
201}
202
203impl Default for ProxyConnectState {
204    fn default() -> Self {
205        Self::new()
206    }
207}
208
209impl TcpConnectionTracker {
210    /// Create a new tracker with the given connection limit.
211    pub fn new(max_tcp_connections: Option<NonZeroUsize>) -> Self {
212        Self {
213            connections: HashMap::new(),
214            connection_keys: HashSet::new(),
215            max_tcp_connections,
216            rejected_connections: 0,
217        }
218    }
219
220    /// Returns `true` if a tracked socket already exists for this exact
221    /// connection (same source AND destination). O(1) via HashSet lookup.
222    pub fn has_socket_for(&self, src: &SocketAddr, dst: &SocketAddr) -> bool {
223        self.connection_keys.contains(&(*src, *dst))
224    }
225
226    /// Create a smoltcp TCP socket for an incoming SYN and register it.
227    ///
228    /// The socket is put into LISTEN state on the destination IP + port so
229    /// smoltcp will complete the three-way handshake when it processes the
230    /// SYN frame. Binding to the specific destination IP (not just port)
231    /// prevents socket dispatch ambiguity when multiple connections target
232    /// different IPs on the same port.
233    ///
234    /// Returns `false` if at `max_tcp_connections` limit.
235    pub fn create_tcp_socket(
236        &mut self,
237        src: SocketAddr,
238        dst: SocketAddr,
239        sockets: &mut SocketSet<'_>,
240    ) -> bool {
241        self.insert_tcp_socket(src, dst, sockets, false)
242    }
243
244    /// Like [`Self::create_tcp_socket`], for a flow egress policy has
245    /// already denied. The handshake completes so the guest's HTTP/HTTPS
246    /// client can be answered with `403 Forbidden`; the dispatcher never
247    /// dials upstream for it.
248    pub fn create_policy_denied_tcp_socket(
249        &mut self,
250        src: SocketAddr,
251        dst: SocketAddr,
252        sockets: &mut SocketSet<'_>,
253    ) -> bool {
254        self.insert_tcp_socket(src, dst, sockets, true)
255    }
256
257    fn insert_tcp_socket(
258        &mut self,
259        src: SocketAddr,
260        dst: SocketAddr,
261        sockets: &mut SocketSet<'_>,
262        policy_denied: bool,
263    ) -> bool {
264        if self
265            .max_tcp_connections
266            .is_some_and(|max| self.connections.len() >= max.get())
267        {
268            // Reclaim completed flows before rejecting a burst. Existing
269            // listeners have already consumed their SYN in the poll loop;
270            // an idle listener here is an invalid or reset handshake.
271            self.cleanup_closed(sockets);
272            if self
273                .max_tcp_connections
274                .is_some_and(|max| self.connections.len() >= max.get())
275            {
276                self.rejected_connections = self.rejected_connections.saturating_add(1);
277                return false;
278            }
279        }
280
281        // Create smoltcp TCP socket with buffers.
282        let rx_buf = tcp::SocketBuffer::new(vec![0u8; TCP_RX_BUF_SIZE]);
283        let tx_buf = tcp::SocketBuffer::new(vec![0u8; TCP_TX_BUF_SIZE]);
284        let mut socket = tcp::Socket::new(rx_buf, tx_buf);
285
286        // Listen on the specific destination IP + port. With any_ip mode,
287        // binding to the IP ensures the correct socket accepts each SYN
288        // when multiple connections target the same port on different IPs.
289        let listen_endpoint = IpListenEndpoint {
290            addr: Some(dst.ip().into()),
291            port: dst.port(),
292        };
293        if socket.listen(listen_endpoint).is_err() {
294            return false;
295        }
296
297        let handle = sockets.add(socket);
298
299        // Create channel pairs for proxy task communication.
300        //
301        // smoltcp → proxy (guest sends data, proxy relays to server):
302        let (to_proxy_tx, to_proxy_rx) = mpsc::channel(CHANNEL_CAPACITY);
303        // proxy → smoltcp (server sends data, proxy relays to guest):
304        let (from_proxy_tx, from_proxy_rx) = mpsc::channel(CHANNEL_CAPACITY);
305
306        self.connection_keys.insert((src, dst));
307        self.connections.insert(
308            handle,
309            Connection {
310                src,
311                dst,
312                to_proxy: Some(to_proxy_tx),
313                from_proxy: from_proxy_rx,
314                proxy_channels: Some(ProxyChannels {
315                    from_smoltcp: to_proxy_rx,
316                    to_smoltcp: from_proxy_tx,
317                }),
318                proxy_spawned: false,
319                proxy_connect: Arc::new(ProxyConnectState::new()),
320                write_buf: None,
321                read_buf: None,
322                deferred_close: DeferredClose::default(),
323                policy_denied,
324            },
325        );
326
327        true
328    }
329
330    /// Earliest pending drain deadline for the network poll loop.
331    pub(crate) fn deferred_close_delay(&self) -> Option<std::time::Duration> {
332        self.connections
333            .values()
334            .filter_map(|conn| conn.deferred_close.poll_delay())
335            .min()
336    }
337
338    /// Relay data between smoltcp sockets and proxy task channels.
339    ///
340    /// For each connection with a spawned proxy:
341    /// - Reads data from the smoltcp socket and sends it to the proxy channel.
342    /// - Receives data from the proxy channel and writes it to the smoltcp socket.
343    pub fn relay_data(&mut self, sockets: &mut SocketSet<'_>) {
344        let mut relay_buf = [0u8; RELAY_BUF_SIZE];
345
346        for (&handle, conn) in &mut self.connections {
347            if !conn.proxy_spawned {
348                continue;
349            }
350
351            let socket = sockets.get_mut::<tcp::Socket>(handle);
352
353            // Already torn down (e.g. abort fired on a previous pass).
354            // Leave it for `cleanup_closed` to evict.
355            if matches!(socket.state(), tcp::State::Closed) {
356                conn.deferred_close = DeferredClose::default();
357                continue;
358            }
359
360            // Detect proxy task exit: when the proxy drops its channel
361            // ends, close the smoltcp socket so the guest gets a FIN.
362            //
363            // If the proxy attempted and failed to reach upstream,
364            // an RST via `abort()` is instead sent so happy-eyeballs
365            // clients fall back to another family instead of committing
366            // to this half-open connection.
367            let proxy_exited = match &conn.to_proxy {
368                Some(to_proxy) => to_proxy.is_closed(),
369                // The guest already half-closed (sender dropped below), so
370                // proxy exit is detected on the other channel instead: the
371                // proxy drops its `to_smoltcp` sender when it returns.
372                None => conn.from_proxy.is_closed(),
373            };
374            if proxy_exited {
375                if matches!(
376                    conn.proxy_connect.status(),
377                    ProxyConnectStatus::UpstreamConnectFailed
378                ) {
379                    tracing::debug!(
380                        src = %conn.src,
381                        dst = %conn.dst,
382                        "upstream connect failed; aborting smoltcp socket (RST to guest)"
383                    );
384                    socket.abort();
385                    continue;
386                }
387                let queued_before = socket.send_queue();
388                write_proxy_data(socket, conn);
389                let written = socket.send_queue() - queued_before;
390                conn.deferred_close
391                    .finish(socket, conn.write_buf.is_some(), written);
392                continue;
393            }
394
395            // smoltcp → proxy: flush read_buf first, then read from socket.
396            if let Some(to_proxy) = &conn.to_proxy {
397                if let Some(pending) = conn.read_buf.take()
398                    && let Err(e) = to_proxy.try_send(pending)
399                {
400                    conn.read_buf = Some(e.into_inner());
401                }
402
403                if conn.read_buf.is_none() {
404                    while socket.can_recv() {
405                        match socket.recv_slice(&mut relay_buf) {
406                            Ok(n) if n > 0 => {
407                                let data = Bytes::copy_from_slice(&relay_buf[..n]);
408                                if let Err(e) = to_proxy.try_send(data) {
409                                    conn.read_buf = Some(e.into_inner());
410                                    break;
411                                }
412                            }
413                            _ => break,
414                        }
415                    }
416                }
417
418                // Guest half-close: the guest sent a FIN (CLOSE_WAIT) and
419                // everything it sent has been relayed. Drop the sender so
420                // the proxy task sees EOF and can shut down the guest →
421                // server direction upstream. The server → guest direction
422                // stays open; the socket is closed once the proxy task
423                // exits (see `proxy_exited` above).
424                if matches!(socket.state(), tcp::State::CloseWait)
425                    && conn.read_buf.is_none()
426                    && !socket.can_recv()
427                {
428                    conn.to_proxy = None;
429                }
430            }
431
432            // proxy → smoltcp: write pending data, then drain channel.
433            write_proxy_data(socket, conn);
434        }
435    }
436
437    /// Collect newly-established connections that need proxy tasks.
438    ///
439    /// Returns a list of [`NewConnection`] structs containing the channel ends
440    /// for the proxy task. The poll loop is responsible for spawning the task.
441    pub fn take_new_connections(&mut self, sockets: &mut SocketSet<'_>) -> Vec<NewConnection> {
442        let mut new = Vec::new();
443
444        for (&handle, conn) in &mut self.connections {
445            if conn.proxy_spawned {
446                continue;
447            }
448
449            let socket = sockets.get::<tcp::Socket>(handle);
450            if matches!(
451                socket.state(),
452                tcp::State::Established | tcp::State::CloseWait
453            ) {
454                conn.proxy_spawned = true;
455
456                if let Some(channels) = conn.proxy_channels.take() {
457                    new.push(NewConnection {
458                        dst: conn.dst,
459                        from_smoltcp: channels.from_smoltcp,
460                        to_smoltcp: channels.to_smoltcp,
461                        proxy_connect: conn.proxy_connect.clone(),
462                        policy_denied: conn.policy_denied,
463                    });
464                }
465            }
466        }
467
468        new
469    }
470
471    /// Record bounded-cardinality diagnostics once per maintenance interval.
472    pub fn trace_stats(&self, sockets: &SocketSet<'_>) {
473        if !tracing::enabled!(target: PROFILING_TARGET, tracing::Level::TRACE) {
474            return;
475        }
476        let closing = self
477            .connections
478            .keys()
479            .filter(|&&handle| {
480                matches!(
481                    sockets.get::<tcp::Socket>(handle).state(),
482                    tcp::State::CloseWait
483                        | tcp::State::FinWait1
484                        | tcp::State::FinWait2
485                        | tcp::State::Closing
486                        | tcp::State::LastAck
487                        | tcp::State::TimeWait
488                )
489            })
490            .count();
491        tracing::trace!(
492            target: PROFILING_TARGET,
493            limit = ?self.max_tcp_connections,
494            tracked = self.connections.len(),
495            closing,
496            rejected_total = self.rejected_connections,
497            socket_buffer_bytes = self.connections.len() * (TCP_RX_BUF_SIZE + TCP_TX_BUF_SIZE),
498            "TCP connection budget"
499        );
500    }
501
502    /// Remove closed connections and their sockets.
503    ///
504    /// Idle listeners represent failed/reset SYNs: this tracker never owns
505    /// persistent listening sockets. Closed sockets with a remote endpoint
506    /// still owe the guest an RST and must survive until smoltcp emits it.
507    /// TIME_WAIT remains intact to reject delayed duplicate segments.
508    pub fn cleanup_closed(&mut self, sockets: &mut SocketSet<'_>) {
509        let keys = &mut self.connection_keys;
510        self.connections.retain(|&handle, conn| {
511            let socket = sockets.get::<tcp::Socket>(handle);
512            if matches!(socket.state(), tcp::State::Closed | tcp::State::Listen)
513                && socket.remote_endpoint().is_none()
514            {
515                keys.remove(&(conn.src, conn.dst));
516                sockets.remove(handle);
517                false
518            } else {
519                true
520            }
521        });
522    }
523}
524
525//--------------------------------------------------------------------------------------------------
526// Functions
527//--------------------------------------------------------------------------------------------------
528
529/// Try to write proxy data to the smoltcp socket.
530fn write_proxy_data(socket: &mut tcp::Socket<'_>, conn: &mut Connection) {
531    // First, try to finish writing any pending partial data.
532    if let Some((data, offset)) = &mut conn.write_buf {
533        if socket.can_send() {
534            match socket.send_slice(&data[*offset..]) {
535                Ok(written) => {
536                    *offset += written;
537                    if *offset >= data.len() {
538                        conn.write_buf = None;
539                    }
540                }
541                Err(_) => return,
542            }
543        } else {
544            return;
545        }
546    }
547
548    // Then drain the channel.
549    while conn.write_buf.is_none() {
550        match conn.from_proxy.try_recv() {
551            Ok(data) => {
552                if socket.can_send() {
553                    match socket.send_slice(&data) {
554                        Ok(written) if written < data.len() => {
555                            conn.write_buf = Some((data, written));
556                        }
557                        Err(_) => {
558                            conn.write_buf = Some((data, 0));
559                        }
560                        _ => {}
561                    }
562                } else {
563                    conn.write_buf = Some((data, 0));
564                }
565            }
566            Err(_) => break,
567        }
568    }
569}
570
571//--------------------------------------------------------------------------------------------------
572// Tests
573//--------------------------------------------------------------------------------------------------
574
575#[cfg(test)]
576mod tests {
577    use super::*;
578
579    #[tokio::test(start_paused = true)]
580    async fn exited_proxy_drains_slow_guest_and_times_out_only_when_stalled() {
581        use crate::tcp::test_support::TestNetwork;
582
583        for stalled in [false, true] {
584            let mut network = TestNetwork::new(true);
585            let mut tracker = TcpConnectionTracker::new(None);
586            assert!(tracker.create_tcp_socket(
587                "10.0.0.1:12345".parse().unwrap(),
588                "10.0.0.2:8099".parse().unwrap(),
589                &mut network.sockets,
590            ));
591            for _ in 0..16 {
592                network.poll();
593            }
594            let mut connections = tracker.take_new_connections(&mut network.sockets);
595            assert_eq!(connections.len(), 1);
596            let conn = connections.remove(0);
597            conn.proxy_connect.mark_connected();
598            let payload: Vec<u8> = (0..262144).map(|i| (i % 251) as u8).collect();
599            for chunk in payload.chunks(16384) {
600                conn.to_smoltcp
601                    .try_send(Bytes::copy_from_slice(chunk))
602                    .unwrap();
603            }
604            drop(conn);
605
606            network
607                .check_drain(|sockets| tracker.relay_data(sockets), &payload, stalled)
608                .await;
609        }
610    }
611
612    #[test]
613    fn omitted_limit_tracks_more_than_the_previous_default() {
614        let mut tracker = TcpConnectionTracker::new(None);
615        let mut sockets = SocketSet::new(Vec::new());
616        let dst = "198.51.100.1:443".parse().unwrap();
617        for port in 10000..10300 {
618            let src = SocketAddr::from(([192, 0, 2, 1], port));
619            assert!(tracker.create_tcp_socket(src, dst, &mut sockets));
620        }
621        assert_eq!(tracker.connections.len(), 300);
622    }
623}