eggress_protocol_http/connect/
client.rs1use tokio::io::{AsyncReadExt, AsyncWriteExt};
2
3use crate::error::HttpError;
4use eggress_core::{BoxStream, TargetAddr, TargetHost};
5
6#[derive(Debug, Clone)]
8pub struct HttpConnectLimits {
9 pub max_status_line: usize,
11 pub max_headers_bytes: usize,
13 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
27pub 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
39pub async fn http_connect(
52 mut stream: BoxStream,
53 target: &TargetAddr,
54 auth: Option<(&str, &str)>,
55 limits: &HttpConnectLimits,
56) -> Result<BoxStream, HttpError> {
57 if let Some((user, pass)) = auth {
59 validate_credentials(user)?;
60 validate_credentials(pass)?;
61 }
62
63 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 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 let response = read_response_head(&mut stream, limits).await?;
88
89 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
117async 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 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 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
163pub 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
191fn 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
221async 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()); }
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 #[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()); 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}