use color_eyre::eyre::{Context, ContextCompat, Result, bail, ensure};
use std::{marker::PhantomData, net::ToSocketAddrs, pin::Pin};
use url::Url;
pub struct Builder<S> {
data: Vec<u8>,
conn: Conn,
_stage: PhantomData<S>,
}
#[derive(Debug, PartialEq, Hash, Clone)]
pub enum Conn {
Tcp {
scheme: ConnScheme,
host: String,
port: Option<u16>,
},
Uds {
path: String,
},
}
#[derive(Debug, Hash, Clone, Copy, PartialEq)]
pub enum ConnScheme {
Http,
Https,
}
pub struct RequestBuilder;
pub struct HeaderStage;
pub type HeaderBuilder = Builder<HeaderStage>;
#[derive(Debug)]
pub struct Request {
pub data: Pin<Box<[u8]>>,
pub conn: Conn,
}
impl Request {
pub fn create_socket(&self) -> Result<(socket2::Socket, socket2::SockAddr)> {
let (socket_addr, domain, socket_type, protocol) = match &self.conn {
Conn::Tcp {
host,
port,
scheme: _,
} => {
let socket_addr = match port {
Some(port) => (host.as_str(), *port).to_socket_addrs(),
None => host.to_socket_addrs(),
};
let mut socket_addr = socket_addr
.wrap_err_with(|| format!("to socket addrs failed for {}:{:?}", host, port))?;
let addr = socket_addr
.next()
.wrap_err_with(|| format!("to socket addrs failed for {}:{:?}", host, port))?;
(
socket2::SockAddr::from(addr),
socket2::Domain::IPV4,
socket2::Type::STREAM,
Some(socket2::Protocol::TCP),
)
}
Conn::Uds { path } => {
let sock_addr = socket2::SockAddr::unix(path).wrap_err_with(|| {
format!("Failed to create Unix socket address for path: {}", path)
})?;
(
sock_addr,
socket2::Domain::UNIX,
socket2::Type::STREAM,
None,
)
}
};
let socket = socket2::Socket::new(domain, socket_type, protocol)
.wrap_err_with(|| format!("Failed to create socket for connection: {:?}", self.conn))?;
socket
.set_nonblocking(true)
.wrap_err("Failed to set socket to nonblocking mode")?;
socket
.set_keepalive(true)
.wrap_err("Failed to set keep‑alive on socket")?;
Ok((socket, socket_addr))
}
}
impl RequestBuilder {
pub fn get(uri: &str) -> Result<HeaderBuilder> {
let url = Url::parse(uri).wrap_err("url parsing")?;
let mut v = Vec::new();
v.extend_from_slice("GET ".as_bytes());
let conn = match url.scheme() {
"http" | "https" => {
v.extend_from_slice(url.path().as_bytes());
if let Some(query) = url.query() {
v.extend_from_slice("?".as_bytes());
v.extend_from_slice(query.as_bytes());
}
v.extend_from_slice(" HTTP/1.1\r\nHost: ".as_bytes());
let host = url.host_str().expect("http always has host").to_string();
v.extend_from_slice(host.as_bytes());
if let Some(port) = url.port()
&& port != 80
{
v.extend_from_slice(format!(":{}", port).as_bytes());
}
v.extend_from_slice("\r\n".as_bytes());
let scheme = match url.scheme() {
"http" => ConnScheme::Http,
"https" => ConnScheme::Https,
_ => unreachable!(),
};
Conn::Tcp {
host,
port: url.port(),
scheme,
}
}
"unix" => {
let socket_path = match url.path().split_once("//") {
None => {
v.extend_from_slice("/".as_bytes());
String::from(url.path())
}
Some((socket_path, resource_path)) => {
v.extend_from_slice("/".as_bytes());
v.extend_from_slice(resource_path.as_bytes());
String::from(socket_path)
}
};
ensure!(!url.has_host(), "uds should not have host");
v.extend_from_slice(" HTTP/1.1\r\nHost: localhost\r\n".as_bytes());
Conn::Uds { path: socket_path }
}
scheme => {
bail!("{} not supported", scheme);
}
};
Ok(HeaderBuilder {
data: v,
conn,
_stage: PhantomData,
})
}
}
impl HeaderBuilder {
fn is_valid_header_name(name: &str) -> bool {
name.chars().all(|c| {
matches!(c,
'!' | '#' | '$' | '%' | '&' | '\'' | '*' | '+' | '-' | '.' |
'^' | '_' | '`' | '|' | '~' |
'0'..='9' | 'a'..='z' | 'A'..='Z'
)
}) && !name.is_empty()
}
fn is_valid_header_value(value: &str) -> bool {
!value.contains('\r') && !value.contains('\n')
}
pub fn append_header(self, key: &str, value: &str) -> Result<Self> {
ensure!(
Self::is_valid_header_name(key),
"{} is not valid header name",
key
);
ensure!(
Self::is_valid_header_value(value),
"{} is not valid header value",
value
);
self.append_header_unchecked(key, value)
}
pub fn append_header_unchecked(mut self, key: &str, value: &str) -> Result<Self> {
self.data.extend_from_slice(key.as_bytes());
self.data.extend_from_slice(": ".as_bytes());
self.data.extend_from_slice(value.as_bytes());
self.data.extend_from_slice("\r\n".as_bytes());
Ok(self)
}
#[allow(dead_code)]
pub fn build_with_body(mut self, body: &[u8]) -> Request {
self.data
.extend_from_slice(format!("Content-Length: {}\r\n\r\n", body.len()).as_bytes());
self.data.extend_from_slice(body);
Request {
data: Pin::new(self.data.into_boxed_slice()),
conn: self.conn,
}
}
pub fn build_without_body(mut self) -> Request {
self.data.extend_from_slice("\r\n".as_bytes());
Request {
data: Pin::new(self.data.into_boxed_slice()),
conn: self.conn,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tcp_http_request() {
let builder = RequestBuilder::get("http://example.com/foo/bar").unwrap();
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with("GET /foo/bar HTTP/1.1\r\nHost: example.com\r\n"));
}
#[test]
fn test_unix_socket_request() {
let sock_path = "/tmp/my.sock";
let builder = RequestBuilder::get(&format!("unix://{}", sock_path)).unwrap();
match &builder.conn {
Conn::Uds { path } => assert_eq!(path, "/tmp/my.sock"),
_ => panic!("Expected UDS connection"),
}
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with("GET / HTTP/1.1\r\nHost: localhost\r\n"));
}
#[test]
fn test_unix_socket_request_resource() {
let sock_path = "/tmp/my.sock";
let resource_path = "/movies/1";
let builder =
RequestBuilder::get(&format!("unix://{}/{}", sock_path, resource_path)).unwrap();
match &builder.conn {
Conn::Uds { path } => assert_eq!(path, "/tmp/my.sock"),
_ => panic!("Expected UDS connection"),
}
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with(&format!(
"GET {} HTTP/1.1\r\nHost: localhost\r\n",
resource_path
)));
}
#[test]
fn test_non_supported_scheme() {
let res = RequestBuilder::get("ftp://example.com/x");
assert!(res.is_err());
}
#[test]
fn test_uds_with_host_error() {
let res = RequestBuilder::get("unix://host/tmp/my.sock");
assert!(res.is_err());
}
#[test]
fn test_tcp_http_request_with_port_default() {
let builder = RequestBuilder::get("http://example.com:80/foo/bar").unwrap();
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with("GET /foo/bar HTTP/1.1\r\nHost: example.com\r\n"));
}
#[test]
fn test_tcp_http_request_with_port() {
let builder = RequestBuilder::get("http://example.com:8080/foo/bar").unwrap();
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with("GET /foo/bar HTTP/1.1\r\nHost: example.com:8080\r\n"));
}
#[test]
fn test_tcp_http_request_root_path() {
let builder = RequestBuilder::get("http://example.com/").unwrap();
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with("GET / HTTP/1.1\r\nHost: example.com\r\n"));
}
#[test]
fn test_unix_socket_request_complex_resource() {
let sock_path = "/tmp/my.sock";
let resource_path = "/movies//1"; let builder =
RequestBuilder::get(&format!("unix://{}/{}", sock_path, resource_path)).unwrap();
match &builder.conn {
Conn::Uds { path } => assert_eq!(path, "/tmp/my.sock"),
_ => panic!("Expected UDS connection"),
}
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with(&format!(
"GET {} HTTP/1.1\r\nHost: localhost\r\n",
resource_path
)));
}
#[test]
fn test_invalid_url() {
let res = RequestBuilder::get("not-a-url");
assert!(res.is_err());
}
#[test]
fn test_append_single_header() {
let request = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("Content-Type", "application/json")
.unwrap()
.build_without_body();
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.contains("Content-Type: application/json\r\n"));
}
#[test]
fn test_append_multiple_headers() {
let request = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("Content-Type", "application/json")
.unwrap()
.append_header("Accept", "application/json")
.unwrap()
.build_without_body();
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.contains("Content-Type: application/json\r\n"));
assert!(req_str.contains("Accept: application/json\r\n"));
}
#[test]
fn test_append_invalid_header_name() {
let res = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("Invalid Name", "value");
assert!(res.is_err());
let res_empty = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("", "value");
assert!(res_empty.is_err());
}
#[test]
fn test_append_invalid_header_value() {
let res = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("Name", "value\r\n");
assert!(res.is_err());
}
#[test]
fn test_build_with_body() {
let body = "{\"key\":\"value\"}";
let request = RequestBuilder::get("http://example.com/foo")
.unwrap()
.build_with_body(body.as_bytes());
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.ends_with(body));
assert!(req_str.contains(&format!("Content-Length: {}\r\n", body.len())));
}
#[test]
fn test_build_with_empty_body() {
let body = "";
let request = RequestBuilder::get("http://example.com/foo")
.unwrap()
.build_with_body(body.as_bytes());
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.ends_with("\r\n\r\n")); assert!(req_str.contains("Content-Length: 0\r\n"));
}
#[test]
fn test_build_without_body() {
let request = RequestBuilder::get("http://example.com/foo")
.unwrap()
.build_without_body();
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.ends_with("\r\n\r\n"));
assert!(!req_str.contains("Content-Length:"));
}
#[test]
fn test_header_name_valid_special_characters() {
let valid_chars =
"!#$%&'*+-.^_`|~0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ";
let request = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header(valid_chars, "test-value")
.unwrap()
.build_without_body();
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.contains(&format!("{}: test-value\r\n", valid_chars)));
}
#[test]
fn test_header_name_invalid_characters() {
let invalid_names = vec![
"name with space",
"name\twith\ttab",
"name@with@at",
"name(with)parens",
"name[with]brackets",
"name{with}braces",
"name\"with\"quotes",
"name\\with\\backslash",
"name/with/slash",
"name:with:colon",
"name;with;semicolon",
"name<with>angles",
"name=with=equals",
"name?with?question",
];
for invalid_name in invalid_names {
let res = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header(invalid_name, "value");
assert!(
res.is_err(),
"Expected error for header name: {}",
invalid_name
);
}
}
#[test]
fn test_header_value_individual_control_characters() {
let res_cr = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("Name", "value\rafter");
assert!(res_cr.is_err());
let res_lf = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("Name", "value\nafter");
assert!(res_lf.is_err());
}
#[test]
fn test_header_value_valid_whitespace() {
let request = RequestBuilder::get("http://example.com/foo")
.unwrap()
.append_header("Name", "value with spaces\tand\ttabs")
.unwrap()
.build_without_body();
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.contains("Name: value with spaces\tand\ttabs\r\n"));
}
#[test]
fn test_url_with_query_parameters() {
let builder = RequestBuilder::get("http://example.com/search?q=test&limit=10").unwrap();
let req_str = String::from_utf8_lossy(&builder.data);
assert!(
req_str.starts_with("GET /search?q=test&limit=10 HTTP/1.1\r\nHost: example.com\r\n")
);
}
#[test]
fn test_url_with_fragment() {
let builder = RequestBuilder::get("http://example.com/page#section1").unwrap();
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with("GET /page HTTP/1.1\r\nHost: example.com\r\n"));
}
#[test]
fn test_url_with_query_and_fragment() {
let builder = RequestBuilder::get("http://example.com/search?q=test#results").unwrap();
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with("GET /search?q=test HTTP/1.1\r\nHost: example.com\r\n"));
}
#[test]
fn test_tcp_connection_values() {
let builder = RequestBuilder::get("http://example.com:8080/foo").unwrap();
match &builder.conn {
Conn::Tcp {
host,
port,
scheme: _,
} => {
assert_eq!(host.as_str(), "example.com");
assert_eq!(*port, Some(8080));
}
_ => panic!("Expected TCP connection"),
}
}
#[test]
fn test_tcp_connection_default_port() {
let builder = RequestBuilder::get("http://example.com/foo").unwrap();
match &builder.conn {
Conn::Tcp { port, .. } => {
assert_eq!(*port, None); }
_ => panic!("Expected TCP connection"),
}
}
#[test]
fn test_unix_socket_edge_case_paths() {
let test_cases = vec![
("unix:///var/run/socket.sock", "/var/run/socket.sock", "/"),
("unix:///tmp/app.sock//api/v1", "/tmp/app.sock", "/api/v1"),
(
"unix:///home/user/.socket//very//long//path",
"/home/user/.socket",
"/very//long//path",
),
];
for (url, expected_socket, expected_resource) in test_cases {
let builder = RequestBuilder::get(url).unwrap();
match &builder.conn {
Conn::Uds { path } => assert_eq!(path, expected_socket),
_ => panic!("Expected UDS connection for {}", url),
}
let req_str = String::from_utf8_lossy(&builder.data);
assert!(req_str.starts_with(&format!(
"GET {} HTTP/1.1\r\nHost: localhost\r\n",
expected_resource
)));
}
}
#[test]
fn test_malformed_urls() {
let malformed_urls = vec![
"http://", "ftp://example.com/file", "://example.com", "http://[invalid-ipv6", ];
for url in malformed_urls {
let res = RequestBuilder::get(url);
assert!(res.is_err(), "Expected error for malformed URL: {}", url);
}
}
#[test]
fn test_complete_request_flow() {
let body = r#"{"user": "test", "action": "login"}"#;
let request = RequestBuilder::get("http://api.example.com:3000/auth")
.unwrap()
.append_header("Content-Type", "application/json")
.unwrap()
.append_header("User-Agent", "concuring/0.1")
.unwrap()
.append_header("Accept", "application/json")
.unwrap()
.build_with_body(body.as_bytes());
let req_str = String::from_utf8_lossy(&request.data);
assert!(req_str.starts_with("GET /auth HTTP/1.1\r\nHost: api.example.com:3000\r\n"));
assert!(req_str.contains("Content-Type: application/json\r\n"));
assert!(req_str.contains("User-Agent: concuring/0.1\r\n"));
assert!(req_str.contains("Accept: application/json\r\n"));
assert!(req_str.contains(&format!("Content-Length: {}\r\n", body.len())));
assert!(req_str.ends_with(body));
assert!(req_str.contains("\r\n\r\n")); }
}