smolvm_protocol/
intercept.rs1use std::io::{self, Read, Write};
31use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
32
33const MAGIC: &[u8; 8] = b"SMOLICPT";
34const VERSION: u8 = 1;
35pub const TOKEN_LEN: usize = 32;
37pub const HEADER_LEN: usize = MAGIC.len() + 1 + TOKEN_LEN + 1 + 2;
39
40pub 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
50pub const VERDICT_CONNECTED: u8 = 0;
52
53pub 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
71pub 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#[derive(Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
86#[serde(deny_unknown_fields)]
87pub struct InterceptEndpoint {
88 pub addr: SocketAddr,
90 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 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 pub fn write_preamble<W: Write>(&self, mut w: W, destination: SocketAddr) -> io::Result<()> {
140 w.write_all(&self.preamble(destination))
141 }
142
143 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 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}