eggress_protocol_http/connect/
server.rs1use std::net::IpAddr;
2
3use tokio::io::{AsyncReadExt, AsyncWriteExt};
4
5use crate::error::HttpError;
6use eggress_core::{BoxStream, TargetAddr, TargetHost};
7
8const MAX_HEAD_SIZE: usize = 32 * 1024;
10
11const MAX_HEADER_LINES: usize = 128;
13
14#[derive(Debug, Clone)]
16pub struct ConnectRequest {
17 pub target: TargetAddr,
18 pub proxy_auth: Option<(String, String)>,
19}
20
21pub async fn handle_connect(
35 mut stream: BoxStream,
36 require_auth: bool,
37 valid_credentials: Option<(&str, &str)>,
38) -> Result<(ConnectRequest, BoxStream), HttpError> {
39 let request = read_connect_request(&mut stream).await?;
40
41 if require_auth {
43 match &request.proxy_auth {
44 Some((user, pass)) => {
45 if let Some((valid_user, valid_pass)) = valid_credentials {
46 use subtle::ConstantTimeEq;
47 let user_ok: bool = user.as_bytes().ct_eq(valid_user.as_bytes()).into();
48 let pass_ok: bool = pass.as_bytes().ct_eq(valid_pass.as_bytes()).into();
49 if !user_ok || !pass_ok {
50 write_error_response(&mut stream, 407, "Proxy Authentication Required")
51 .await?;
52 return Err(HttpError::AuthRequired);
53 }
54 } else {
55 write_error_response(&mut stream, 407, "Proxy Authentication Required").await?;
56 return Err(HttpError::AuthRequired);
57 }
58 }
59 None => {
60 write_error_response(&mut stream, 407, "Proxy Authentication Required").await?;
61 return Err(HttpError::AuthRequired);
62 }
63 }
64 }
65
66 stream
68 .write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
69 .await?;
70 stream.flush().await?;
71
72 Ok((request, stream))
73}
74
75async fn read_connect_request(stream: &mut BoxStream) -> Result<ConnectRequest, HttpError> {
77 let mut head_buf = Vec::with_capacity(1024);
78 let mut temp = [0u8; 1];
79 let mut header_count = 0;
80
81 loop {
82 if head_buf.len() >= MAX_HEAD_SIZE {
83 return Err(HttpError::HeaderTooLarge);
84 }
85
86 let n = stream.read(&mut temp).await?;
87 if n == 0 {
88 return Err(HttpError::MalformedRequest(
89 "unexpected EOF reading request".into(),
90 ));
91 }
92
93 head_buf.push(temp[0]);
94
95 if head_buf.len() >= 4 {
97 let len = head_buf.len();
98 if &head_buf[len - 4..] == b"\r\n\r\n" {
99 break;
100 }
101 if head_buf.len() >= 2 && &head_buf[len - 2..] == b"\r\n" {
103 header_count += 1;
104 if header_count > MAX_HEADER_LINES {
105 return Err(HttpError::TooManyHeaders);
106 }
107 }
108 }
109 }
110
111 let head_str = String::from_utf8_lossy(&head_buf);
112 let mut lines = head_str.split("\r\n");
113
114 let request_line = lines
116 .next()
117 .ok_or_else(|| HttpError::MalformedRequest("empty request".into()))?;
118
119 let parts: Vec<&str> = request_line.split_whitespace().collect();
120 if parts.len() != 3 {
121 return Err(HttpError::MalformedRequest(format!(
122 "expected 3 parts in request line, got {}",
123 parts.len()
124 )));
125 }
126
127 if parts[0] != "CONNECT" {
128 return Err(HttpError::MalformedRequest(format!(
129 "expected CONNECT method, got {}",
130 parts[0]
131 )));
132 }
133
134 if parts[2] != "HTTP/1.1" && parts[2] != "HTTP/1.0" {
135 return Err(HttpError::UnsupportedVersion(parts[2].to_string()));
136 }
137
138 let authority = parts[1];
140 let target = parse_authority(authority)?;
141
142 let mut proxy_auth = None;
144 for line in lines {
145 if line.is_empty() {
146 break;
147 }
148 if let Some((name, value)) = parse_header_line(line) {
149 if name.eq_ignore_ascii_case("Proxy-Authorization") {
150 proxy_auth = parse_basic_auth(&value);
151 }
152 }
153 }
154
155 Ok(ConnectRequest { target, proxy_auth })
156}
157
158pub fn parse_authority(authority: &str) -> Result<TargetAddr, HttpError> {
163 if authority.starts_with('[') {
165 let bracket_end = authority.find(']').ok_or_else(|| {
166 HttpError::TargetParseError("unclosed bracket in IPv6 address".into())
167 })?;
168
169 let ip_str = &authority[1..bracket_end];
170 let ip: IpAddr = ip_str
171 .parse()
172 .map_err(|e| HttpError::TargetParseError(format!("invalid IPv6 address: {}", e)))?;
173
174 let port_str = authority
175 .get(bracket_end + 2..)
176 .ok_or_else(|| HttpError::TargetParseError("missing port after IPv6 address".into()))?;
177
178 if !authority
179 .as_bytes()
180 .get(bracket_end + 1)
181 .is_some_and(|&b| b == b':')
182 {
183 return Err(HttpError::TargetParseError(
184 "expected ':' between IPv6 address and port".into(),
185 ));
186 }
187
188 let port: u16 = port_str
189 .parse()
190 .map_err(|e| HttpError::TargetParseError(format!("invalid port: {}", e)))?;
191
192 return Ok(TargetAddr {
193 host: TargetHost::Ip(ip),
194 port,
195 });
196 }
197
198 let colon_pos = authority
201 .rfind(':')
202 .ok_or_else(|| HttpError::TargetParseError("missing port in authority".into()))?;
203
204 let host_str = &authority[..colon_pos];
205 let port_str = &authority[colon_pos + 1..];
206
207 let port: u16 = port_str
208 .parse()
209 .map_err(|e| HttpError::TargetParseError(format!("invalid port: {}", e)))?;
210
211 if let Ok(ip) = host_str.parse::<IpAddr>() {
213 return Ok(TargetAddr {
214 host: TargetHost::Ip(ip),
215 port,
216 });
217 }
218
219 if host_str.is_empty() {
221 return Err(HttpError::TargetParseError("empty host".into()));
222 }
223
224 Ok(TargetAddr {
225 host: TargetHost::Domain(host_str.to_string()),
226 port,
227 })
228}
229
230pub fn parse_header_line(line: &str) -> Option<(String, String)> {
234 let colon_pos = line.find(':')?;
235 let name = line[..colon_pos].trim().to_string();
236 let value = line[colon_pos + 1..].trim().to_string();
237 Some((name, value))
238}
239
240pub fn parse_basic_auth(value: &str) -> Option<(String, String)> {
244 let value = value.trim();
245 if !value.starts_with("Basic ") {
246 return None;
247 }
248
249 let encoded = &value[6..];
250 let decoded = base64_decode(encoded)?;
251 let decoded_str = String::from_utf8(decoded).ok()?;
252 let colon_pos = decoded_str.find(':')?;
253 let username = decoded_str[..colon_pos].to_string();
254 let password = decoded_str[colon_pos + 1..].to_string();
255 Some((username, password))
256}
257
258fn base64_decode(input: &str) -> Option<Vec<u8>> {
260 const TABLE: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
261
262 let input = input.trim_end_matches('=');
263 let input_bytes = input.as_bytes();
264
265 let mut result = Vec::with_capacity(input_bytes.len() * 3 / 4);
266 let mut buf: u32 = 0;
267 let mut bits: u32 = 0;
268
269 for &byte in input_bytes {
270 let val = TABLE.iter().position(|&b| b == byte)? as u32;
271 buf = (buf << 6) | val;
272 bits += 6;
273 if bits >= 8 {
274 bits -= 8;
275 result.push((buf >> bits) as u8);
276 }
277 }
278
279 Some(result)
280}
281
282async fn write_error_response(
284 stream: &mut BoxStream,
285 status: u16,
286 reason: &str,
287) -> Result<(), HttpError> {
288 let response = format!(
289 "HTTP/1.1 {} {}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
290 status, reason
291 );
292 stream.write_all(response.as_bytes()).await?;
293 stream.flush().await?;
294 Ok(())
295}
296
297#[cfg(test)]
298mod tests {
299 use super::*;
300
301 #[test]
302 fn test_parse_authority_ipv4() {
303 let target = parse_authority("192.168.1.1:8080").unwrap();
304 assert_eq!(
305 target,
306 TargetAddr {
307 host: TargetHost::Ip("192.168.1.1".parse().unwrap()),
308 port: 8080,
309 }
310 );
311 }
312
313 #[test]
314 fn test_parse_authority_ipv6() {
315 let target = parse_authority("[::1]:443").unwrap();
316 assert_eq!(
317 target,
318 TargetAddr {
319 host: TargetHost::Ip("::1".parse().unwrap()),
320 port: 443,
321 }
322 );
323 }
324
325 #[test]
326 fn test_parse_authority_domain() {
327 let target = parse_authority("example.com:443").unwrap();
328 assert_eq!(
329 target,
330 TargetAddr {
331 host: TargetHost::Domain("example.com".to_string()),
332 port: 443,
333 }
334 );
335 }
336
337 #[test]
338 fn test_parse_authority_missing_port() {
339 assert!(parse_authority("example.com").is_err());
340 }
341
342 #[test]
343 fn test_parse_header_line() {
344 let (name, value) = parse_header_line("Host: example.com").unwrap();
345 assert_eq!(name, "Host");
346 assert_eq!(value, "example.com");
347 }
348
349 #[test]
350 fn test_parse_basic_auth() {
351 let result = parse_basic_auth("Basic dXNlcjpwYXNz").unwrap();
353 assert_eq!(result, ("user".to_string(), "pass".to_string()));
354 }
355
356 #[test]
357 fn test_parse_basic_auth_no_prefix() {
358 assert!(parse_basic_auth("Bearer token").is_none());
359 }
360
361 #[test]
362 fn test_base64_decode() {
363 let decoded = base64_decode("dGVzdA==").unwrap();
364 assert_eq!(decoded, b"test");
365 }
366
367 #[test]
368 fn test_max_head_size_enforced() {
369 assert_eq!(MAX_HEAD_SIZE, 32 * 1024);
370 assert_eq!(MAX_HEADER_LINES, 128);
371 }
372
373 #[test]
374 fn test_parse_authority_empty_host() {
375 assert!(parse_authority(":80").is_err());
376 }
377
378 #[test]
379 fn test_parse_authority_empty_string() {
380 assert!(parse_authority("").is_err());
381 }
382
383 #[test]
384 fn test_parse_authority_no_colon() {
385 assert!(parse_authority("example.com").is_err());
386 }
387
388 #[test]
389 fn test_parse_header_line_no_colon() {
390 assert!(parse_header_line("no-colon-here").is_none());
391 }
392
393 #[test]
394 fn test_parse_header_line_empty() {
395 assert!(parse_header_line("").is_none());
396 }
397
398 #[test]
399 fn test_parse_basic_auth_not_basic() {
400 assert!(parse_basic_auth("Bearer token123").is_none());
401 }
402
403 #[test]
404 fn test_parse_basic_auth_invalid_base64() {
405 assert!(parse_basic_auth("Basic !!!invalid!!!").is_none());
406 }
407
408 #[tokio::test]
409 async fn test_head_too_large_rejected() {
410 use tokio::io::{AsyncReadExt, AsyncWriteExt};
411
412 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
413 let addr = listener.local_addr().unwrap();
414 let jh = tokio::spawn(async move {
415 let (mut stream, _) = listener.accept().await.unwrap();
416 let mut payload = b"CONNECT example.com:443 HTTP/1.1\r\n".to_vec();
418 let header_line = b"X-Pad: AAAAAAAAAAAAAAAAAAAAAAAAAAAAA\r\n";
420 while payload.len() < MAX_HEAD_SIZE + header_line.len() {
421 payload.extend_from_slice(header_line);
422 }
423 payload.extend_from_slice(b"\r\n");
424 let _ = stream.write_all(&payload).await;
425 tokio::time::sleep(std::time::Duration::from_millis(200)).await;
427 });
428
429 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
430 let mut buf = vec![0u8; 4096];
431 let _ =
433 tokio::time::timeout(std::time::Duration::from_secs(2), stream.read(&mut buf)).await;
434 jh.abort();
435 }
436
437 #[tokio::test]
438 async fn test_too_many_header_lines_rejected() {
439 use tokio::io::{AsyncReadExt, AsyncWriteExt};
440
441 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
442 let addr = listener.local_addr().unwrap();
443 let jh = tokio::spawn(async move {
444 let (mut stream, _) = listener.accept().await.unwrap();
445 let mut payload = b"CONNECT example.com:443 HTTP/1.1\r\n".to_vec();
447 for i in 0..=MAX_HEADER_LINES + 1 {
448 payload.extend_from_slice(format!("X-Header-{i}: value\r\n").as_bytes());
449 }
450 payload.extend_from_slice(b"\r\n");
451 let _ = stream.write_all(&payload).await;
452 tokio::time::sleep(std::time::Duration::from_millis(200)).await;
453 });
454
455 let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
456 let mut buf = vec![0u8; 4096];
457 let _ =
458 tokio::time::timeout(std::time::Duration::from_secs(2), stream.read(&mut buf)).await;
459 jh.abort();
460 }
461}