use std::net::IpAddr;
use base64::Engine;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use crate::error::HttpError;
use eggress_core::{BoxStream, TargetAddr, TargetHost};
const MAX_HEAD_SIZE: usize = 32 * 1024;
const MAX_HEADER_LINES: usize = 128;
#[derive(Debug, Clone)]
pub struct ConnectRequest {
pub target: TargetAddr,
pub proxy_auth: Option<(String, String)>,
}
pub async fn handle_connect(
stream: BoxStream,
require_auth: bool,
valid_credentials: Option<(&str, &str)>,
) -> Result<(ConnectRequest, BoxStream), HttpError> {
let mut stream: BoxStream = Box::new(tokio::io::BufReader::new(stream));
let request = read_connect_request(&mut stream).await?;
if require_auth {
match &request.proxy_auth {
Some((user, pass)) => {
if let Some((valid_user, valid_pass)) = valid_credentials {
use subtle::ConstantTimeEq;
let user_ok: bool = user.as_bytes().ct_eq(valid_user.as_bytes()).into();
let pass_ok: bool = pass.as_bytes().ct_eq(valid_pass.as_bytes()).into();
if !user_ok || !pass_ok {
write_error_response(&mut stream, 407, "Proxy Authentication Required")
.await?;
return Err(HttpError::AuthRequired);
}
} else {
write_error_response(&mut stream, 407, "Proxy Authentication Required").await?;
return Err(HttpError::AuthRequired);
}
}
None => {
write_error_response(&mut stream, 407, "Proxy Authentication Required").await?;
return Err(HttpError::AuthRequired);
}
}
}
stream
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await?;
stream.flush().await?;
Ok((request, stream))
}
async fn read_connect_request(stream: &mut BoxStream) -> Result<ConnectRequest, HttpError> {
let mut head_buf = Vec::with_capacity(1024);
let mut temp = [0u8; 1];
let mut header_count = 0;
let mut saw_request_line = false;
loop {
if head_buf.len() >= MAX_HEAD_SIZE {
return Err(HttpError::HeaderTooLarge);
}
let n = stream.read(&mut temp).await?;
if n == 0 {
return Err(HttpError::MalformedRequest(
"unexpected EOF reading request".into(),
));
}
head_buf.push(temp[0]);
if head_buf.len() >= 4 {
let len = head_buf.len();
if &head_buf[len - 4..] == b"\r\n\r\n" {
break;
}
if head_buf.len() >= 2 && &head_buf[len - 2..] == b"\r\n" {
if saw_request_line {
header_count += 1;
} else {
saw_request_line = true;
}
if header_count > MAX_HEADER_LINES {
return Err(HttpError::TooManyHeaders);
}
}
}
}
let head_str = String::from_utf8_lossy(&head_buf);
let mut lines = head_str.split("\r\n");
let request_line = lines
.next()
.ok_or_else(|| HttpError::MalformedRequest("empty request".into()))?;
let parts: Vec<&str> = request_line.split_whitespace().collect();
if parts.len() != 3 {
return Err(HttpError::MalformedRequest(format!(
"expected 3 parts in request line, got {}",
parts.len()
)));
}
if parts[0] != "CONNECT" {
return Err(HttpError::MalformedRequest(format!(
"expected CONNECT method, got {}",
parts[0]
)));
}
if parts[2] != "HTTP/1.1" && parts[2] != "HTTP/1.0" {
return Err(HttpError::UnsupportedVersion(parts[2].to_string()));
}
let authority = parts[1];
let target = parse_authority(authority)?;
let mut proxy_auth = None;
for line in lines {
if line.is_empty() {
break;
}
if let Some((name, value)) = parse_header_line(line) {
if name.eq_ignore_ascii_case("Proxy-Authorization") {
proxy_auth = parse_basic_auth(&value);
}
}
}
Ok(ConnectRequest { target, proxy_auth })
}
pub fn parse_authority(authority: &str) -> Result<TargetAddr, HttpError> {
if authority.starts_with('[') {
let bracket_end = authority.find(']').ok_or_else(|| {
HttpError::TargetParseError("unclosed bracket in IPv6 address".into())
})?;
let ip_str = &authority[1..bracket_end];
let ip: IpAddr = ip_str
.parse()
.map_err(|e| HttpError::TargetParseError(format!("invalid IPv6 address: {}", e)))?;
let port_str = authority
.get(bracket_end + 2..)
.ok_or_else(|| HttpError::TargetParseError("missing port after IPv6 address".into()))?;
if authority
.as_bytes()
.get(bracket_end + 1)
.is_none_or(|&b| b != b':')
{
return Err(HttpError::TargetParseError(
"expected ':' between IPv6 address and port".into(),
));
}
let port: u16 = port_str
.parse()
.map_err(|e| HttpError::TargetParseError(format!("invalid port: {}", e)))?;
return Ok(TargetAddr {
host: TargetHost::Ip(ip),
port,
});
}
const DEFAULT_CONNECT_PORT: u16 = 443;
let (host_str, port) = match authority.rfind(':') {
Some(colon_pos) => {
let port: u16 = authority[colon_pos + 1..]
.parse()
.map_err(|e| HttpError::TargetParseError(format!("invalid port: {}", e)))?;
(&authority[..colon_pos], port)
}
None => (authority, DEFAULT_CONNECT_PORT),
};
if let Ok(ip) = host_str.parse::<IpAddr>() {
return Ok(TargetAddr {
host: TargetHost::Ip(ip),
port,
});
}
if host_str.is_empty() {
return Err(HttpError::TargetParseError("empty host".into()));
}
Ok(TargetAddr {
host: TargetHost::Domain(host_str.to_string()),
port,
})
}
pub fn parse_header_line(line: &str) -> Option<(String, String)> {
let colon_pos = line.find(':')?;
let name = line[..colon_pos].trim().to_string();
let value = line[colon_pos + 1..].trim().to_string();
Some((name, value))
}
pub fn parse_basic_auth(value: &str) -> Option<(String, String)> {
let value = value.trim();
if !value.starts_with("Basic ") {
return None;
}
let encoded = &value[6..];
let decoded = base64::engine::general_purpose::STANDARD
.decode(encoded)
.ok()?;
let decoded_str = String::from_utf8(decoded).ok()?;
let colon_pos = decoded_str.find(':')?;
let username = decoded_str[..colon_pos].to_string();
let password = decoded_str[colon_pos + 1..].to_string();
Some((username, password))
}
async fn write_error_response(
stream: &mut BoxStream,
status: u16,
reason: &str,
) -> Result<(), HttpError> {
let response = format!(
"HTTP/1.1 {} {}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
status, reason
);
stream.write_all(response.as_bytes()).await?;
stream.flush().await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_authority_ipv4() {
let target = parse_authority("192.168.1.1:8080").unwrap();
assert_eq!(
target,
TargetAddr {
host: TargetHost::Ip("192.168.1.1".parse().unwrap()),
port: 8080,
}
);
}
#[test]
fn test_parse_authority_ipv6() {
let target = parse_authority("[::1]:443").unwrap();
assert_eq!(
target,
TargetAddr {
host: TargetHost::Ip("::1".parse().unwrap()),
port: 443,
}
);
}
#[test]
fn test_parse_authority_domain() {
let target = parse_authority("example.com:443").unwrap();
assert_eq!(
target,
TargetAddr {
host: TargetHost::Domain("example.com".to_string()),
port: 443,
}
);
}
#[test]
fn test_parse_authority_missing_port_implies_default() {
let target = parse_authority("example.com").unwrap();
assert_eq!(
target,
TargetAddr {
host: TargetHost::Domain("example.com".to_string()),
port: 443,
}
);
let target = parse_authority("192.168.1.1").unwrap();
assert_eq!(
target,
TargetAddr {
host: TargetHost::Ip("192.168.1.1".parse().unwrap()),
port: 443,
}
);
}
#[test]
fn test_parse_header_line() {
let (name, value) = parse_header_line("Host: example.com").unwrap();
assert_eq!(name, "Host");
assert_eq!(value, "example.com");
}
#[test]
fn test_parse_basic_auth() {
let result = parse_basic_auth("Basic dXNlcjpwYXNz").unwrap();
assert_eq!(result, ("user".to_string(), "pass".to_string()));
}
#[test]
fn test_parse_basic_auth_no_prefix() {
assert!(parse_basic_auth("Bearer token").is_none());
}
#[test]
fn test_base64_decode() {
let decoded = base64::engine::general_purpose::STANDARD
.decode("dGVzdA==")
.unwrap();
assert_eq!(decoded, b"test");
}
#[test]
fn test_max_head_size_enforced() {
assert_eq!(MAX_HEAD_SIZE, 32 * 1024);
assert_eq!(MAX_HEADER_LINES, 128);
}
#[test]
fn test_parse_authority_empty_string() {
assert!(parse_authority("").is_err());
}
#[test]
fn test_parse_authority_empty_host_with_port() {
assert!(parse_authority(":80").is_err());
assert!(parse_authority("").is_err());
}
#[test]
fn test_parse_header_line_no_colon() {
assert!(parse_header_line("no-colon-here").is_none());
}
#[test]
fn test_parse_header_line_empty() {
assert!(parse_header_line("").is_none());
}
#[test]
fn test_parse_basic_auth_not_basic() {
assert!(parse_basic_auth("Bearer token123").is_none());
}
#[test]
fn test_parse_basic_auth_invalid_base64() {
assert!(parse_basic_auth("Basic !!!invalid!!!").is_none());
}
#[tokio::test]
async fn test_head_too_large_rejected() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let jh = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut payload = b"CONNECT example.com:443 HTTP/1.1\r\n".to_vec();
let header_line = b"X-Pad: AAAAAAAAAAAAAAAAAAAAAAAAAAAAA\r\n";
while payload.len() < MAX_HEAD_SIZE + header_line.len() {
payload.extend_from_slice(header_line);
}
payload.extend_from_slice(b"\r\n");
let _ = stream.write_all(&payload).await;
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
});
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let mut buf = vec![0u8; 4096];
let _ =
tokio::time::timeout(std::time::Duration::from_secs(2), stream.read(&mut buf)).await;
jh.abort();
}
#[tokio::test]
async fn test_too_many_header_lines_rejected() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let jh = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut payload = b"CONNECT example.com:443 HTTP/1.1\r\n".to_vec();
for i in 0..=MAX_HEADER_LINES + 1 {
payload.extend_from_slice(format!("X-Header-{i}: value\r\n").as_bytes());
}
payload.extend_from_slice(b"\r\n");
let _ = stream.write_all(&payload).await;
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
});
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let mut buf = vec![0u8; 4096];
let _ =
tokio::time::timeout(std::time::Duration::from_secs(2), stream.read(&mut buf)).await;
jh.abort();
}
}