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, Debug, PartialEq, Eq)]
86pub struct InterceptEndpoint {
87    /// Loopback listener owned by the interceptor.
88    pub addr: SocketAddr,
89    /// Secret shared between the interceptor and the machine's backend only.
90    pub token: [u8; TOKEN_LEN],
91}
92
93impl InterceptEndpoint {
94    /// Encode the preamble for one redirected flow.
95    pub fn preamble(&self, destination: SocketAddr) -> Vec<u8> {
96        let mut out = Vec::with_capacity(HEADER_LEN + 16);
97        out.extend_from_slice(MAGIC);
98        out.push(VERSION);
99        out.extend_from_slice(&self.token);
100        match destination.ip() {
101            IpAddr::V4(_) => out.push(4),
102            IpAddr::V6(_) => out.push(6),
103        }
104        out.extend_from_slice(&destination.port().to_be_bytes());
105        match destination.ip() {
106            IpAddr::V4(ip) => out.extend_from_slice(&ip.octets()),
107            IpAddr::V6(ip) => out.extend_from_slice(&ip.octets()),
108        }
109        out
110    }
111
112    /// Write the preamble for `destination` to a freshly connected stream.
113    pub fn write_preamble<W: Write>(&self, mut w: W, destination: SocketAddr) -> io::Result<()> {
114        w.write_all(&self.preamble(destination))
115    }
116
117    /// Read and authenticate a preamble from an accepted connection.
118    ///
119    /// Returns the destination the guest dialed. A wrong magic, version or
120    /// token is reported as `InvalidData` after consuming the fixed header so
121    /// the caller can simply drop the connection; the token comparison does not
122    /// short-circuit on the first differing byte.
123    pub fn read_preamble<R: Read>(&self, mut r: R) -> io::Result<SocketAddr> {
124        let mut header = [0u8; HEADER_LEN];
125        r.read_exact(&mut header)?;
126        let (magic, rest) = header.split_at(MAGIC.len());
127        let (version, rest) = rest.split_first().expect("fixed header");
128        let (token, rest) = rest.split_at(TOKEN_LEN);
129        let (family, port) = rest.split_first().expect("fixed header");
130        let port = u16::from_be_bytes([port[0], port[1]]);
131
132        let mut mismatch = (magic != MAGIC) as u8 | (*version != VERSION) as u8;
133        for (a, b) in token.iter().zip(self.token.iter()) {
134            mismatch |= a ^ b;
135        }
136        if mismatch != 0 {
137            return Err(io::Error::new(
138                io::ErrorKind::InvalidData,
139                "intercept preamble rejected",
140            ));
141        }
142        let ip = match family {
143            4 => {
144                let mut octets = [0u8; 4];
145                r.read_exact(&mut octets)?;
146                IpAddr::V4(Ipv4Addr::from(octets))
147            }
148            6 => {
149                let mut octets = [0u8; 16];
150                r.read_exact(&mut octets)?;
151                IpAddr::V6(Ipv6Addr::from(octets))
152            }
153            _ => {
154                return Err(io::Error::new(
155                    io::ErrorKind::InvalidData,
156                    "intercept preamble has an unknown address family",
157                ))
158            }
159        };
160        Ok(SocketAddr::new(ip, port))
161    }
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    fn endpoint(token: u8) -> InterceptEndpoint {
169        InterceptEndpoint {
170            addr: "127.0.0.1:1".parse().unwrap(),
171            token: [token; TOKEN_LEN],
172        }
173    }
174
175    #[test]
176    fn round_trips_v4_and_v6_destinations() {
177        let ep = endpoint(7);
178        for dst in [
179            "93.184.216.34:443",
180            "[2606:2800:220:1:248:1893:25c8:1946]:8443",
181        ] {
182            let dst: SocketAddr = dst.parse().unwrap();
183            let bytes = ep.preamble(dst);
184            let mut cursor = std::io::Cursor::new(bytes);
185            assert_eq!(ep.read_preamble(&mut cursor).unwrap(), dst);
186            assert_eq!(cursor.position() as usize, cursor.get_ref().len());
187        }
188    }
189
190    #[test]
191    fn rejects_wrong_token_and_magic() {
192        let dst: SocketAddr = "93.184.216.34:443".parse().unwrap();
193        let bytes = endpoint(1).preamble(dst);
194        let err = endpoint(2).read_preamble(&bytes[..]).unwrap_err();
195        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
196
197        let mut bad_magic = endpoint(1).preamble(dst);
198        bad_magic[0] = b'X';
199        let err = endpoint(1).read_preamble(&bad_magic[..]).unwrap_err();
200        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
201    }
202
203    #[test]
204    fn preamble_len_follows_the_family_byte() {
205        let dst4: SocketAddr = "93.184.216.34:443".parse().unwrap();
206        let dst6: SocketAddr = "[2001:db8::1]:443".parse().unwrap();
207        let p4 = endpoint(1).preamble(dst4);
208        let p6 = endpoint(1).preamble(dst6);
209        assert_eq!(preamble_len(&p4[..HEADER_LEN]), Some(p4.len()));
210        assert_eq!(preamble_len(&p6[..HEADER_LEN]), Some(p6.len()));
211        // Too short to hold the family byte, or a family we do not speak.
212        assert_eq!(preamble_len(&p4[..HEADER_LEN - 4]), None);
213        let mut bad_family = p4.clone();
214        bad_family[HEADER_LEN - 3] = 5;
215        assert_eq!(preamble_len(&bad_family[..HEADER_LEN]), None);
216    }
217
218    #[test]
219    fn verdicts_carry_the_connect_outcome() {
220        let unreachable: io::Result<()> = Err(io::ErrorKind::HostUnreachable.into());
221        assert_eq!(connect_verdict(&Ok::<(), io::Error>(())), VERDICT_CONNECTED);
222        assert_eq!(connect_verdict(&unreachable), 113);
223        read_verdict(&[VERDICT_CONNECTED][..]).unwrap();
224        assert!(read_verdict(&[113u8][..]).is_err());
225        assert_eq!(
226            read_verdict(&[][..]).unwrap_err().kind(),
227            io::ErrorKind::UnexpectedEof
228        );
229    }
230
231    #[test]
232    fn short_reads_surface_as_eof() {
233        let dst: SocketAddr = "93.184.216.34:443".parse().unwrap();
234        let bytes = endpoint(1).preamble(dst);
235        let err = endpoint(1)
236            .read_preamble(&bytes[..bytes.len() - 1])
237            .unwrap_err();
238        assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
239    }
240}