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