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, Debug, PartialEq, Eq)]
86pub struct InterceptEndpoint {
87 pub addr: SocketAddr,
89 pub token: [u8; TOKEN_LEN],
91}
92
93impl InterceptEndpoint {
94 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 pub fn write_preamble<W: Write>(&self, mut w: W, destination: SocketAddr) -> io::Result<()> {
114 w.write_all(&self.preamble(destination))
115 }
116
117 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 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}