eggress_protocol_http/connect/
client.rs1use 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#[derive(Debug, Clone)]
12pub struct HttpConnectLimits {
13 pub max_status_line: usize,
15 pub max_headers_bytes: usize,
17 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
31pub 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
43pub async fn http_connect(
56 stream: BoxStream,
57 target: &TargetAddr,
58 auth: Option<(&str, &str)>,
59 limits: &HttpConnectLimits,
60) -> Result<BoxStream, HttpError> {
61 let mut stream: BoxStream = Box::new(BufferedStream {
65 reader: BufReader::new(stream),
66 });
67
68 if let Some((user, pass)) = auth {
70 validate_credentials(user)?;
71 validate_credentials(pass)?;
72 }
73
74 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 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 let response = read_response_head(&mut stream, limits).await?;
99
100 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
145async 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 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 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
191pub 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()); }
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 #[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()); 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}