Skip to main content

smolvm_protocol/
intercept.rs

1//! Host-side stream interception handshake.
2//!
3//! When a machine carries a credential policy, its network backend redirects
4//! selected guest TCP flows (HTTPS by default) to an interceptor listening on
5//! host loopback instead of dialing the destination directly. The redirecting
6//! side — the virtio-net relay in `smolvm-network`, or libkrun's TSI muxer —
7//! prefixes the redirected byte stream with a fixed-size preamble that names
8//! the destination the guest actually asked for and proves the connection came
9//! from the machine's own backend rather than from an arbitrary host process
10//! that found the loopback port.
11//!
12//! Wire layout (all integers big-endian):
13//!
14//! ```text
15//! magic    8 bytes  "SMOLICPT"
16//! version  1 byte   1
17//! token   32 bytes  per-machine secret shared with the interceptor
18//! family   1 byte   4 or 6
19//! port     2 bytes  destination port
20//! address  4 or 16  destination IP
21//! ```
22//!
23//! The interceptor answers with one byte before any payload flows: `0` once
24//! it has connected to the destination, otherwise the Linux errno of that
25//! connect. libkrun's TSI muxer reports it to the guest as the result of the
26//! guest's own connect, so a guest with an unreachable IPv6 route still falls
27//! back to IPv4 exactly as it would without interception. Guest payload then
28//! follows; the guest never sees the preamble or the verdict.
29
30use std::io::{self, Read, Write};
31use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
32
33const MAGIC: &[u8; 8] = b"SMOLICPT";
34const VERSION: u8 = 1;
35/// Bytes of shared secret carried in every preamble.
36pub const TOKEN_LEN: usize = 32;
37/// Bytes of preamble before the destination address.
38pub const HEADER_LEN: usize = MAGIC.len() + 1 + TOKEN_LEN + 1 + 2;
39
40/// Total preamble length implied by an already-read fixed header, or `None`
41/// if the family byte is invalid. Lets an async reader size its second read.
42pub fn preamble_len(header: &[u8]) -> Option<usize> {
43    match header.get(HEADER_LEN - 3)? {
44        4 => Some(HEADER_LEN + 4),
45        6 => Some(HEADER_LEN + 16),
46        _ => None,
47    }
48}
49
50/// Verdict byte for a destination the interceptor reached.
51pub const VERDICT_CONNECTED: u8 = 0;
52
53/// Verdict byte for the interceptor's connect to the real destination: `0`, or
54/// the Linux errno the guest should see (guests are Linux whatever the host).
55pub fn connect_verdict<T>(result: &io::Result<T>) -> u8 {
56    const ENETUNREACH: u8 = 101;
57    const ETIMEDOUT: u8 = 110;
58    const ECONNREFUSED: u8 = 111;
59    const EHOSTUNREACH: u8 = 113;
60    match result {
61        Ok(_) => VERDICT_CONNECTED,
62        Err(e) => match e.kind() {
63            io::ErrorKind::NetworkUnreachable => ENETUNREACH,
64            io::ErrorKind::HostUnreachable => EHOSTUNREACH,
65            io::ErrorKind::TimedOut => ETIMEDOUT,
66            _ => ECONNREFUSED,
67        },
68    }
69}
70
71/// Read the interceptor's verdict; an error carries the errno it reported.
72pub fn read_verdict<R: Read>(mut r: R) -> io::Result<()> {
73    let mut verdict = [0u8; 1];
74    r.read_exact(&mut verdict)?;
75    match verdict[0] {
76        VERDICT_CONNECTED => Ok(()),
77        errno => Err(io::Error::new(
78            io::ErrorKind::ConnectionRefused,
79            format!("interceptor could not reach the destination (errno {errno})"),
80        )),
81    }
82}
83
84/// Where a backend redirects intercepted flows, and the secret it must present.
85#[derive(Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
86#[serde(deny_unknown_fields)]
87pub struct InterceptEndpoint {
88    /// Loopback listener owned by the interceptor.
89    pub addr: SocketAddr,
90    /// Secret shared between the interceptor and the machine's backend only.
91    pub token: [u8; TOKEN_LEN],
92}
93
94impl std::fmt::Debug for InterceptEndpoint {
95    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
96        f.debug_struct("InterceptEndpoint")
97            .field("addr", &self.addr)
98            .field("token", &"<redacted>")
99            .finish()
100    }
101}
102
103#[cfg(test)]
104mod debug_tests {
105    use super::InterceptEndpoint;
106
107    #[test]
108    fn interceptor_token_is_redacted() {
109        let endpoint = InterceptEndpoint {
110            addr: "127.0.0.1:9000".parse().unwrap(),
111            token: [42; 32],
112        };
113        let rendered = format!("{endpoint:?}");
114        assert!(rendered.contains("<redacted>"));
115        assert!(!rendered.contains("42"));
116    }
117}
118
119impl InterceptEndpoint {
120    /// Encode the preamble for one redirected flow.
121    pub fn preamble(&self, destination: SocketAddr) -> Vec<u8> {
122        let mut out = Vec::with_capacity(HEADER_LEN + 16);
123        out.extend_from_slice(MAGIC);
124        out.push(VERSION);
125        out.extend_from_slice(&self.token);
126        match destination.ip() {
127            IpAddr::V4(_) => out.push(4),
128            IpAddr::V6(_) => out.push(6),
129        }
130        out.extend_from_slice(&destination.port().to_be_bytes());
131        match destination.ip() {
132            IpAddr::V4(ip) => out.extend_from_slice(&ip.octets()),
133            IpAddr::V6(ip) => out.extend_from_slice(&ip.octets()),
134        }
135        out
136    }
137
138    /// Write the preamble for `destination` to a freshly connected stream.
139    pub fn write_preamble<W: Write>(&self, mut w: W, destination: SocketAddr) -> io::Result<()> {
140        w.write_all(&self.preamble(destination))
141    }
142
143    /// Read and authenticate a preamble from an accepted connection.
144    ///
145    /// Returns the destination the guest dialed. A wrong magic, version or
146    /// token is reported as `InvalidData` after consuming the fixed header so
147    /// the caller can simply drop the connection; the token comparison does not
148    /// short-circuit on the first differing byte.
149    pub fn read_preamble<R: Read>(&self, mut r: R) -> io::Result<SocketAddr> {
150        let mut header = [0u8; HEADER_LEN];
151        r.read_exact(&mut header)?;
152        let (magic, rest) = header.split_at(MAGIC.len());
153        let (version, rest) = rest.split_first().expect("fixed header");
154        let (token, rest) = rest.split_at(TOKEN_LEN);
155        let (family, port) = rest.split_first().expect("fixed header");
156        let port = u16::from_be_bytes([port[0], port[1]]);
157
158        let mut mismatch = (magic != MAGIC) as u8 | (*version != VERSION) as u8;
159        for (a, b) in token.iter().zip(self.token.iter()) {
160            mismatch |= a ^ b;
161        }
162        if mismatch != 0 {
163            return Err(io::Error::new(
164                io::ErrorKind::InvalidData,
165                "intercept preamble rejected",
166            ));
167        }
168        let ip = match family {
169            4 => {
170                let mut octets = [0u8; 4];
171                r.read_exact(&mut octets)?;
172                IpAddr::V4(Ipv4Addr::from(octets))
173            }
174            6 => {
175                let mut octets = [0u8; 16];
176                r.read_exact(&mut octets)?;
177                IpAddr::V6(Ipv6Addr::from(octets))
178            }
179            _ => {
180                return Err(io::Error::new(
181                    io::ErrorKind::InvalidData,
182                    "intercept preamble has an unknown address family",
183                ))
184            }
185        };
186        Ok(SocketAddr::new(ip, port))
187    }
188}
189
190#[cfg(test)]
191mod tests {
192    use super::*;
193
194    fn endpoint(token: u8) -> InterceptEndpoint {
195        InterceptEndpoint {
196            addr: "127.0.0.1:1".parse().unwrap(),
197            token: [token; TOKEN_LEN],
198        }
199    }
200
201    #[test]
202    fn round_trips_v4_and_v6_destinations() {
203        let ep = endpoint(7);
204        for dst in [
205            "93.184.216.34:443",
206            "[2606:2800:220:1:248:1893:25c8:1946]:8443",
207        ] {
208            let dst: SocketAddr = dst.parse().unwrap();
209            let bytes = ep.preamble(dst);
210            let mut cursor = std::io::Cursor::new(bytes);
211            assert_eq!(ep.read_preamble(&mut cursor).unwrap(), dst);
212            assert_eq!(cursor.position() as usize, cursor.get_ref().len());
213        }
214    }
215
216    #[test]
217    fn rejects_wrong_token_and_magic() {
218        let dst: SocketAddr = "93.184.216.34:443".parse().unwrap();
219        let bytes = endpoint(1).preamble(dst);
220        let err = endpoint(2).read_preamble(&bytes[..]).unwrap_err();
221        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
222
223        let mut bad_magic = endpoint(1).preamble(dst);
224        bad_magic[0] = b'X';
225        let err = endpoint(1).read_preamble(&bad_magic[..]).unwrap_err();
226        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
227    }
228
229    #[test]
230    fn preamble_len_follows_the_family_byte() {
231        let dst4: SocketAddr = "93.184.216.34:443".parse().unwrap();
232        let dst6: SocketAddr = "[2001:db8::1]:443".parse().unwrap();
233        let p4 = endpoint(1).preamble(dst4);
234        let p6 = endpoint(1).preamble(dst6);
235        assert_eq!(preamble_len(&p4[..HEADER_LEN]), Some(p4.len()));
236        assert_eq!(preamble_len(&p6[..HEADER_LEN]), Some(p6.len()));
237        // Too short to hold the family byte, or a family we do not speak.
238        assert_eq!(preamble_len(&p4[..HEADER_LEN - 4]), None);
239        let mut bad_family = p4.clone();
240        bad_family[HEADER_LEN - 3] = 5;
241        assert_eq!(preamble_len(&bad_family[..HEADER_LEN]), None);
242    }
243
244    #[test]
245    fn verdicts_carry_the_connect_outcome() {
246        let unreachable: io::Result<()> = Err(io::ErrorKind::HostUnreachable.into());
247        assert_eq!(connect_verdict(&Ok::<(), io::Error>(())), VERDICT_CONNECTED);
248        assert_eq!(connect_verdict(&unreachable), 113);
249        read_verdict(&[VERDICT_CONNECTED][..]).unwrap();
250        assert!(read_verdict(&[113u8][..]).is_err());
251        assert_eq!(
252            read_verdict(&[][..]).unwrap_err().kind(),
253            io::ErrorKind::UnexpectedEof
254        );
255    }
256
257    #[test]
258    fn short_reads_surface_as_eof() {
259        let dst: SocketAddr = "93.184.216.34:443".parse().unwrap();
260        let bytes = endpoint(1).preamble(dst);
261        let err = endpoint(1)
262            .read_preamble(&bytes[..bytes.len() - 1])
263            .unwrap_err();
264        assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
265    }
266}