Skip to main content

eggress_protocol_http/connect/
server.rs

1use std::net::IpAddr;
2
3use tokio::io::{AsyncReadExt, AsyncWriteExt};
4
5use crate::error::HttpError;
6use eggress_core::{BoxStream, TargetAddr, TargetHost};
7
8/// Maximum size for the HTTP request head (request line + headers).
9const MAX_HEAD_SIZE: usize = 32 * 1024;
10
11/// Maximum number of header lines.
12const MAX_HEADER_LINES: usize = 128;
13
14/// Parsed CONNECT request.
15#[derive(Debug, Clone)]
16pub struct ConnectRequest {
17    pub target: TargetAddr,
18    pub proxy_auth: Option<(String, String)>,
19}
20
21/// Handle an HTTP CONNECT request from a client stream.
22///
23/// Parses the CONNECT request, validates it, and returns the stream
24/// ready for bidirectional forwarding after sending a 200 response.
25///
26/// # Arguments
27/// * `stream` - The client stream to read the CONNECT request from
28/// * `require_auth` - Whether proxy authentication is required
29/// * `valid_credentials` - Valid (username, password) pair for auth validation
30///
31/// # Returns
32/// The parsed CONNECT request and the stream (with any bytes after the
33/// request head preserved).
34pub 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    // Validate authentication if required
42    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    // Send success response
67    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
75/// Read and parse an HTTP CONNECT request from the stream.
76async 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        // Check for end of headers (\r\n\r\n)
96        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            // Also count individual \r\n for header line limits
102            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    // Parse request line
115    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    // Parse authority (host:port)
139    let authority = parts[1];
140    let target = parse_authority(authority)?;
141
142    // Parse headers
143    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
158/// Parse an authority-form target (host:port).
159/// Parse an authority string (`host:port` or `[ipv6]:port`) into a [`TargetAddr`].
160///
161/// Exposed for fuzzing. Returns [`HttpError::TargetParseError`] on malformed input.
162pub fn parse_authority(authority: &str) -> Result<TargetAddr, HttpError> {
163    // Handle IPv6 bracketed addresses: [::1]:port
164    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    // Handle IPv4 or domain
199    // Find the last ':' to split host and port
200    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    // Try to parse as IP first
212    if let Ok(ip) = host_str.parse::<IpAddr>() {
213        return Ok(TargetAddr {
214            host: TargetHost::Ip(ip),
215            port,
216        });
217    }
218
219    // Otherwise treat as domain
220    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
230/// Parse a header line into (name, value).
231///
232/// Exposed for fuzzing. Returns `None` if the line lacks a colon.
233pub 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
240/// Parse Basic authentication from a Proxy-Authorization header value.
241///
242/// Exposed for fuzzing. Returns `None` if the value is not a Basic auth header.
243pub 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
258/// Simple base64 decoder (no-std compatible, no external dependency).
259fn 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
282/// Write an HTTP error response.
283async 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        // "user:pass" base64 encoded is "dXNlcjpwYXNz"
352        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            // Send a request line followed by headers that exceed MAX_HEAD_SIZE
417            let mut payload = b"CONNECT example.com:443 HTTP/1.1\r\n".to_vec();
418            // Add headers until we exceed the limit
419            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            // Keep connection alive briefly
426            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        // The server should reject or the client should see an error
432        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            // Send CONNECT with more than MAX_HEADER_LINES (128) header lines
446            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}