Skip to main content

eggress_protocol_http/connect/
client.rs

1use tokio::io::{AsyncReadExt, AsyncWriteExt};
2
3use crate::error::HttpError;
4use eggress_core::{BoxStream, TargetAddr, TargetHost};
5
6/// Configuration limits for HTTP CONNECT response parsing.
7#[derive(Debug, Clone)]
8pub struct HttpConnectLimits {
9    /// Maximum length of the status line (e.g., "HTTP/1.1 200 OK\r\n").
10    pub max_status_line: usize,
11    /// Maximum total bytes for response headers.
12    pub max_headers_bytes: usize,
13    /// Maximum number of header lines (excluding the status line).
14    pub max_header_count: usize,
15}
16
17impl Default for HttpConnectLimits {
18    fn default() -> Self {
19        Self {
20            max_status_line: 1024,
21            max_headers_bytes: 32_768,
22            max_header_count: 100,
23        }
24    }
25}
26
27/// Validate that a credential string contains no control characters.
28///
29/// Control characters are bytes < 0x20 (Space) or 0x7F (DEL).
30pub fn validate_credentials(value: &str) -> Result<(), HttpError> {
31    for byte in value.bytes() {
32        if byte < 0x20 || byte == 0x7F {
33            return Err(HttpError::InvalidCredentials);
34        }
35    }
36    Ok(())
37}
38
39/// Send an HTTP CONNECT request to an upstream proxy and return the
40/// upgraded stream on success.
41///
42/// # Arguments
43/// * `stream` - The stream to the upstream proxy
44/// * `target` - The target address to connect to
45/// * `auth` - Optional (username, password) for Proxy-Authorization
46/// * `limits` - Parsing limits for the response
47///
48/// # Returns
49/// The stream after receiving a 2xx response, ready for bidirectional
50/// forwarding.
51pub async fn http_connect(
52    mut stream: BoxStream,
53    target: &TargetAddr,
54    auth: Option<(&str, &str)>,
55    limits: &HttpConnectLimits,
56) -> Result<BoxStream, HttpError> {
57    // Validate credentials before sending anything
58    if let Some((user, pass)) = auth {
59        validate_credentials(user)?;
60        validate_credentials(pass)?;
61    }
62
63    // Build CONNECT request
64    let host_header = match &target.host {
65        TargetHost::Ip(ip) => format!("{}", ip),
66        TargetHost::Domain(domain) => domain.clone(),
67    };
68
69    let mut request = format!(
70        "CONNECT {}:{} HTTP/1.1\r\nHost: {}:{}\r\n",
71        host_header, target.port, host_header, target.port
72    );
73
74    // Add Proxy-Authorization if provided
75    if let Some((user, pass)) = auth {
76        let credentials = format!("{}:{}", user, pass);
77        let encoded = base64_encode(credentials.as_bytes());
78        request.push_str(&format!("Proxy-Authorization: Basic {}\r\n", encoded));
79    }
80
81    request.push_str("\r\n");
82
83    stream.write_all(request.as_bytes()).await?;
84    stream.flush().await?;
85
86    // Read response
87    let response = read_response_head(&mut stream, limits).await?;
88
89    // Parse status code
90    let status = parse_status_code(&response, limits)?;
91
92    match status {
93        200..=299 => Ok(stream),
94        407 => {
95            let _ = write_error_response(&mut stream, 407, "Proxy Authentication Required").await;
96            Err(HttpError::AuthRequired)
97        }
98        403 => {
99            let _ = write_error_response(&mut stream, 403, "Forbidden").await;
100            Err(HttpError::AuthFailed)
101        }
102        502 => {
103            let _ = write_error_response(&mut stream, 502, "Bad Gateway").await;
104            Err(HttpError::BadGateway)
105        }
106        504 => {
107            let _ = write_error_response(&mut stream, 504, "Gateway Timeout").await;
108            Err(HttpError::GatewayTimeout)
109        }
110        code => {
111            let _ = write_error_response(&mut stream, code, "Upstream Error").await;
112            Err(HttpError::UnexpectedStatus(code))
113        }
114    }
115}
116
117/// Read the HTTP response head (status line + headers) from the stream.
118async fn read_response_head(
119    stream: &mut BoxStream,
120    limits: &HttpConnectLimits,
121) -> Result<String, HttpError> {
122    let mut head_buf = Vec::with_capacity(1024);
123    let mut temp = [0u8; 1];
124    let mut header_count: usize = 0;
125    let mut last_was_cr = false;
126
127    loop {
128        if head_buf.len() >= limits.max_headers_bytes {
129            return Err(HttpError::HeaderTooLarge);
130        }
131
132        let n = stream.read(&mut temp).await?;
133        if n == 0 {
134            return Err(HttpError::MalformedResponse(
135                "unexpected EOF reading response".into(),
136            ));
137        }
138
139        head_buf.push(temp[0]);
140
141        // Count header lines (each \r\n after status line is a header)
142        if temp[0] == b'\n' && last_was_cr {
143            header_count += 1;
144            if header_count > limits.max_header_count {
145                return Err(HttpError::TooManyHeaders);
146            }
147        }
148        last_was_cr = temp[0] == b'\r';
149
150        // Check for end of headers
151        if head_buf.len() >= 4 {
152            let len = head_buf.len();
153            if &head_buf[len - 4..] == b"\r\n\r\n" {
154                break;
155            }
156        }
157    }
158
159    String::from_utf8(head_buf)
160        .map_err(|e| HttpError::MalformedResponse(format!("invalid UTF-8: {}", e)))
161}
162
163/// Parse the HTTP status code from a response head string.
164///
165/// Exposed for fuzzing; takes the full response head (status line + headers)
166/// and returns the numeric status code from the first whitespace-separated
167/// token after the HTTP version.
168pub fn parse_status_code(response: &str, limits: &HttpConnectLimits) -> Result<u16, HttpError> {
169    let first_line = response
170        .lines()
171        .next()
172        .ok_or_else(|| HttpError::MalformedResponse("empty response".into()))?;
173
174    if first_line.len() > limits.max_status_line {
175        return Err(HttpError::MalformedResponse("status line too long".into()));
176    }
177
178    let parts: Vec<&str> = first_line.split_whitespace().collect();
179    if parts.len() < 2 {
180        return Err(HttpError::MalformedResponse(format!(
181            "invalid status line: {}",
182            first_line
183        )));
184    }
185
186    parts[1]
187        .parse::<u16>()
188        .map_err(|e| HttpError::MalformedResponse(format!("invalid status code: {}", e)))
189}
190
191/// Simple base64 encoder (no-std compatible, no external dependency).
192fn base64_encode(input: &[u8]) -> String {
193    const TABLE: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
194
195    let mut result = String::with_capacity(input.len().div_ceil(3) * 4);
196
197    for chunk in input.chunks(3) {
198        let b0 = chunk[0] as u32;
199        let b1 = if chunk.len() > 1 { chunk[1] as u32 } else { 0 };
200        let b2 = if chunk.len() > 2 { chunk[2] as u32 } else { 0 };
201
202        let triple = (b0 << 16) | (b1 << 8) | b2;
203
204        result.push(TABLE[((triple >> 18) & 0x3F) as usize] as char);
205        result.push(TABLE[((triple >> 12) & 0x3F) as usize] as char);
206        if chunk.len() > 1 {
207            result.push(TABLE[((triple >> 6) & 0x3F) as usize] as char);
208        } else {
209            result.push('=');
210        }
211        if chunk.len() > 2 {
212            result.push(TABLE[(triple & 0x3F) as usize] as char);
213        } else {
214            result.push('=');
215        }
216    }
217
218    result
219}
220
221/// Write an HTTP error response.
222async fn write_error_response(
223    stream: &mut BoxStream,
224    status: u16,
225    reason: &str,
226) -> Result<(), HttpError> {
227    let response = format!(
228        "HTTP/1.1 {} {}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
229        status, reason
230    );
231    stream.write_all(response.as_bytes()).await?;
232    stream.flush().await?;
233    Ok(())
234}
235
236#[cfg(test)]
237mod tests {
238    use super::*;
239
240    #[test]
241    fn test_base64_encode() {
242        assert_eq!(base64_encode(b"test"), "dGVzdA==");
243        assert_eq!(base64_encode(b"hello"), "aGVsbG8=");
244        assert_eq!(base64_encode(b"user:pass"), "dXNlcjpwYXNz");
245    }
246
247    #[test]
248    fn test_parse_status_code() {
249        let limits = HttpConnectLimits::default();
250        assert_eq!(
251            parse_status_code("HTTP/1.1 200 Connection Established\r\n", &limits).unwrap(),
252            200
253        );
254        assert_eq!(
255            parse_status_code("HTTP/1.1 407 Proxy Authentication Required\r\n", &limits).unwrap(),
256            407
257        );
258    }
259
260    #[test]
261    fn test_parse_status_code_invalid() {
262        let limits = HttpConnectLimits::default();
263        assert!(parse_status_code("HTTP/1.1", &limits).is_err());
264        assert!(parse_status_code("HTTP/1.1 abc\r\n", &limits).is_err());
265    }
266
267    #[test]
268    fn test_parse_status_code_too_long() {
269        let limits = HttpConnectLimits {
270            max_status_line: 10,
271            ..Default::default()
272        };
273        assert!(parse_status_code("HTTP/1.1 200 OK\r\n", &limits).is_err());
274    }
275
276    #[test]
277    fn test_validate_credentials_rejects_control_chars() {
278        assert!(validate_credentials("user\x00name").is_err());
279        assert!(validate_credentials("user\x1Fname").is_err());
280        assert!(validate_credentials("user\x7Fname").is_err());
281        assert!(validate_credentials("\x01").is_err());
282        assert!(validate_credentials("\x09").is_err()); // TAB
283    }
284
285    #[test]
286    fn test_validate_credentials_accepts_normal() {
287        assert!(validate_credentials("user").is_ok());
288        assert!(validate_credentials("user name").is_ok());
289        assert!(validate_credentials("p@ss:word!").is_ok());
290        assert!(validate_credentials("a]b[c").is_ok());
291    }
292
293    #[test]
294    fn test_http_connect_limits_defaults() {
295        let limits = HttpConnectLimits::default();
296        assert_eq!(limits.max_status_line, 1024);
297        assert_eq!(limits.max_headers_bytes, 32_768);
298        assert_eq!(limits.max_header_count, 100);
299    }
300
301    #[test]
302    fn test_parse_status_code_empty_response() {
303        let limits = HttpConnectLimits::default();
304        assert!(parse_status_code("", &limits).is_err());
305    }
306
307    #[test]
308    fn test_parse_status_code_whitespace_only() {
309        let limits = HttpConnectLimits::default();
310        assert!(parse_status_code("   ", &limits).is_err());
311    }
312
313    // ===== Synthetic server integration tests =====
314
315    #[tokio::test]
316    async fn test_connect_200_success() {
317        use crate::connect::test_server::{ProxyMode, TestProxyServer};
318
319        let server = TestProxyServer::start(ProxyMode::Success).await;
320        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
321        let boxed: BoxStream = Box::new(stream);
322        let target = TargetAddr {
323            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
324            port: 80,
325        };
326        let result = http_connect(boxed, &target, None, &HttpConnectLimits::default()).await;
327        assert!(result.is_ok());
328        server.stop().await;
329    }
330
331    #[tokio::test]
332    async fn test_connect_407_auth_required() {
333        use crate::connect::test_server::{ProxyMode, TestProxyServer};
334
335        let server = TestProxyServer::start(ProxyMode::AuthRequired).await;
336        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
337        let boxed: BoxStream = Box::new(stream);
338        let target = TargetAddr {
339            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
340            port: 80,
341        };
342        let result = http_connect(boxed, &target, None, &HttpConnectLimits::default()).await;
343        assert!(matches!(result, Err(HttpError::AuthRequired)));
344        server.stop().await;
345    }
346
347    #[tokio::test]
348    async fn test_connect_403_forbidden() {
349        use crate::connect::test_server::{ProxyMode, TestProxyServer};
350
351        let server = TestProxyServer::start(ProxyMode::Forbidden).await;
352        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
353        let boxed: BoxStream = Box::new(stream);
354        let target = TargetAddr {
355            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
356            port: 80,
357        };
358        let result = http_connect(boxed, &target, None, &HttpConnectLimits::default()).await;
359        assert!(matches!(result, Err(HttpError::AuthFailed)));
360        server.stop().await;
361    }
362
363    #[tokio::test]
364    async fn test_connect_malformed_status() {
365        use crate::connect::test_server::{ProxyMode, TestProxyServer};
366
367        let server = TestProxyServer::start(ProxyMode::MalformedStatus).await;
368        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
369        let boxed: BoxStream = Box::new(stream);
370        let target = TargetAddr {
371            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
372            port: 80,
373        };
374        let result = http_connect(boxed, &target, None, &HttpConnectLimits::default()).await;
375        assert!(matches!(result, Err(HttpError::MalformedResponse(_))));
376        server.stop().await;
377    }
378
379    #[tokio::test]
380    async fn test_connect_slow_response_timeout() {
381        use crate::connect::test_server::{ProxyMode, TestProxyServer};
382
383        let server =
384            TestProxyServer::start(ProxyMode::SlowResponse(std::time::Duration::from_secs(10)))
385                .await;
386        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
387        let boxed: BoxStream = Box::new(stream);
388        let target = TargetAddr {
389            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
390            port: 80,
391        };
392        let result = tokio::time::timeout(
393            std::time::Duration::from_millis(200),
394            http_connect(boxed, &target, None, &HttpConnectLimits::default()),
395        )
396        .await;
397        assert!(result.is_err()); // timeout
398        server.stop().await;
399    }
400
401    #[tokio::test]
402    async fn test_connect_basic_auth_success() {
403        use crate::connect::test_server::{ProxyMode, TestProxyServer};
404
405        let server = TestProxyServer::start(ProxyMode::Success).await;
406        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
407        let boxed: BoxStream = Box::new(stream);
408        let target = TargetAddr {
409            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
410            port: 80,
411        };
412        let result = http_connect(
413            boxed,
414            &target,
415            Some(("user", "pass")),
416            &HttpConnectLimits::default(),
417        )
418        .await;
419        assert!(result.is_ok());
420        server.stop().await;
421    }
422
423    #[tokio::test]
424    async fn test_connect_basic_auth_wrong() {
425        use crate::connect::test_server::{ProxyMode, TestProxyServer};
426
427        let server = TestProxyServer::start(ProxyMode::AuthRequired).await;
428        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
429        let boxed: BoxStream = Box::new(stream);
430        let target = TargetAddr {
431            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
432            port: 80,
433        };
434        let result = http_connect(
435            boxed,
436            &target,
437            Some(("user", "wrong")),
438            &HttpConnectLimits::default(),
439        )
440        .await;
441        assert!(matches!(result, Err(HttpError::AuthRequired)));
442        server.stop().await;
443    }
444
445    #[tokio::test]
446    async fn test_connect_credentials_with_control_chars_rejected() {
447        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
448        let addr = listener.local_addr().unwrap();
449        let jh = tokio::spawn(async move {
450            let _ = listener.accept().await;
451        });
452        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
453        let boxed: BoxStream = Box::new(stream);
454        let target = TargetAddr {
455            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
456            port: 80,
457        };
458        let result = http_connect(
459            boxed,
460            &target,
461            Some(("user\x00", "pass")),
462            &HttpConnectLimits::default(),
463        )
464        .await;
465        assert!(matches!(result, Err(HttpError::InvalidCredentials)));
466        jh.abort();
467    }
468
469    #[tokio::test]
470    async fn test_connect_headers_too_large() {
471        use crate::connect::test_server::{ProxyMode, TestProxyServer};
472
473        let server = TestProxyServer::start(ProxyMode::HeadersTooLarge).await;
474        let stream = tokio::net::TcpStream::connect(server.addr).await.unwrap();
475        let boxed: BoxStream = Box::new(stream);
476        let target = TargetAddr {
477            host: TargetHost::Ip("127.0.0.1".parse().unwrap()),
478            port: 80,
479        };
480        let result = http_connect(boxed, &target, None, &HttpConnectLimits::default()).await;
481        assert!(matches!(
482            result,
483            Err(HttpError::HeaderTooLarge | HttpError::TooManyHeaders)
484        ));
485        server.stop().await;
486    }
487}