Skip to main content

eggress_protocol_http/connect/
client.rs

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