use std::io::{BufRead, BufReader, Write};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use crate::request::Method;
use crate::Request;
fn drain_request_headers(reader: &mut BufReader<TcpStream>) -> Option<String> {
let mut request_line = String::new();
match reader.read_line(&mut request_line) {
Ok(0) | Err(_) => return None,
Ok(_) => {}
}
loop {
let mut header = String::new();
match reader.read_line(&mut header) {
Ok(0) | Err(_) => return None,
Ok(_) if header == "\r\n" => break,
Ok(_) => {}
}
}
Some(request_line)
}
pub(crate) fn get(url: &str) -> Request {
Request {
name: None,
method: Method::Get,
url: url.to_string(),
headers: Vec::new(),
query: Vec::new(),
body: None,
json: None,
body_file: None,
form: Vec::new(),
multipart: Vec::new(),
auth: None,
assertions: None,
pre_request: None,
post_request: None,
capture: None,
retry: None,
}
}
pub(crate) struct CountingServer {
addr: SocketAddr,
connections: Arc<AtomicUsize>,
requests: Arc<AtomicUsize>,
}
impl CountingServer {
pub(crate) fn start() -> Self {
let listener = TcpListener::bind("127.0.0.1:0").expect("an ephemeral port is free");
let addr = listener.local_addr().expect("the listener has an address");
let connections = Arc::new(AtomicUsize::new(0));
let requests = Arc::new(AtomicUsize::new(0));
let (server_connections, server_requests) = (connections.clone(), requests.clone());
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(stream) = stream else { continue };
server_connections.fetch_add(1, Ordering::SeqCst);
let mut writer = stream.try_clone().expect("the socket clones");
let mut reader = BufReader::new(stream);
loop {
let mut line = String::new();
match reader.read_line(&mut line) {
Ok(0) | Err(_) => break,
Ok(_) => {}
}
loop {
let mut header = String::new();
match reader.read_line(&mut header) {
Ok(0) | Err(_) => return,
Ok(_) if header == "\r\n" => break,
Ok(_) => {}
}
}
server_requests.fetch_add(1, Ordering::SeqCst);
if writer
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok")
.is_err()
{
break;
}
let _ = writer.flush();
}
}
});
Self {
addr,
connections,
requests,
}
}
pub(crate) fn url(&self) -> String {
format!("http://{}/", self.addr)
}
pub(crate) fn connections(&self) -> usize {
self.connections.load(Ordering::SeqCst)
}
pub(crate) fn requests(&self) -> usize {
self.requests.load(Ordering::SeqCst)
}
}
pub(crate) enum Stall {
BeforeResponding,
MidBody,
}
pub(crate) fn start_stalling_server(stall: Stall, delay: Duration) -> SocketAddr {
let listener = TcpListener::bind("127.0.0.1:0").expect("an ephemeral port is free");
let addr = listener.local_addr().expect("the listener has an address");
std::thread::spawn(move || {
let Ok(stream) = listener.accept().map(|(s, _)| s) else {
return;
};
let mut writer = stream.try_clone().expect("the socket clones");
let mut reader = BufReader::new(stream);
if drain_request_headers(&mut reader).is_none() {
return;
}
if matches!(stall, Stall::MidBody) {
let _ = writer.write_all(
b"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 11\r\n\r\n",
);
let _ = writer.flush();
}
std::thread::sleep(delay);
let _ = writer.write_all(b"too late");
});
addr
}
pub(crate) fn start_route_server(routes: Vec<(&'static str, Vec<u8>)>) -> SocketAddr {
let listener = TcpListener::bind("127.0.0.1:0").expect("an ephemeral port is free");
let addr = listener.local_addr().expect("the listener has an address");
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(stream) = stream else { continue };
let mut writer = stream.try_clone().expect("the socket clones");
let mut reader = BufReader::new(stream);
loop {
let Some(request_line) = drain_request_headers(&mut reader) else {
return;
};
let path = request_line
.split_whitespace()
.nth(1)
.unwrap_or("/")
.to_string();
let response = routes
.iter()
.find(|(route, _)| *route == path)
.map(|(_, body)| body.clone())
.unwrap_or_else(|| {
b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n".to_vec()
});
if writer.write_all(&response).is_err() {
return;
}
let _ = writer.flush();
}
}
});
addr
}
pub(crate) fn redirect_response(status: u16, reason: &str, location: &str) -> Vec<u8> {
format!("HTTP/1.1 {status} {reason}\r\nLocation: {location}\r\nContent-Length: 0\r\n\r\n")
.into_bytes()
}
pub(crate) fn redirect_with_cookie_response(
status: u16,
reason: &str,
location: &str,
cookie: &str,
) -> Vec<u8> {
format!(
"HTTP/1.1 {status} {reason}\r\nLocation: {location}\r\nSet-Cookie: {cookie}\r\nContent-Length: 0\r\n\r\n"
)
.into_bytes()
}
pub(crate) fn ok_response(body: &str) -> Vec<u8> {
format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: {}\r\n\r\n{body}",
body.len()
)
.into_bytes()
}
pub(crate) fn set_cookie_response(cookie: &str, body: &str) -> Vec<u8> {
format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nSet-Cookie: {cookie}\r\nContent-Length: {}\r\n\r\n{body}",
body.len()
)
.into_bytes()
}
pub(crate) fn ok_bytes(content_type: &str, body: &[u8]) -> Vec<u8> {
let mut response = format!(
"HTTP/1.1 200 OK\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\n\r\n",
body.len()
)
.into_bytes();
response.extend_from_slice(body);
response
}
pub(crate) fn start_self_signed_tls_server() -> SocketAddr {
let _ = rustls::crypto::ring::default_provider().install_default();
let certified = rcgen::generate_simple_self_signed(["127.0.0.1".to_string()])
.expect("a self-signed certificate for 127.0.0.1 generates");
let cert_der = certified.cert.der().clone();
let key_der =
rustls::pki_types::PrivateKeyDer::Pkcs8(certified.key_pair.serialize_der().into());
let server_config = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(vec![cert_der], key_der)
.expect("the freshly generated cert and key are valid together");
let server_config = Arc::new(server_config);
let listener = TcpListener::bind("127.0.0.1:0").expect("an ephemeral port is free");
let addr = listener.local_addr().expect("the listener has an address");
std::thread::spawn(move || {
let Ok((mut sock, _)) = listener.accept() else {
return;
};
let Ok(mut conn) = rustls::ServerConnection::new(server_config) else {
return;
};
let mut tls = rustls::Stream::new(&mut conn, &mut sock);
let mut reader = BufReader::new(&mut tls);
let mut request_line = String::new();
if reader.read_line(&mut request_line).unwrap_or(0) == 0 {
return;
}
loop {
let mut header = String::new();
match reader.read_line(&mut header) {
Ok(0) | Err(_) => return,
Ok(_) if header == "\r\n" => break,
Ok(_) => {}
}
}
let _ = tls.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok");
});
addr
}
pub(crate) fn start_mutual_tls_server() -> (SocketAddr, String, String) {
let _ = rustls::crypto::ring::default_provider().install_default();
let server_certified = rcgen::generate_simple_self_signed(["127.0.0.1".to_string()])
.expect("a self-signed certificate for 127.0.0.1 generates");
let server_cert_der = server_certified.cert.der().clone();
let server_key_der =
rustls::pki_types::PrivateKeyDer::Pkcs8(server_certified.key_pair.serialize_der().into());
let mut ca_params =
rcgen::CertificateParams::new(Vec::new()).expect("no subject alt names cannot fail");
ca_params.is_ca = rcgen::IsCa::Ca(rcgen::BasicConstraints::Unconstrained);
ca_params
.key_usages
.push(rcgen::KeyUsagePurpose::KeyCertSign);
ca_params
.key_usages
.push(rcgen::KeyUsagePurpose::DigitalSignature);
ca_params.key_usages.push(rcgen::KeyUsagePurpose::CrlSign);
let ca_key = rcgen::KeyPair::generate().expect("key generation does not fail");
let ca_cert = ca_params
.self_signed(&ca_key)
.expect("a self-signed CA certificate generates");
let mut client_params =
rcgen::CertificateParams::new(Vec::new()).expect("no subject alt names cannot fail");
client_params
.key_usages
.push(rcgen::KeyUsagePurpose::DigitalSignature);
client_params
.extended_key_usages
.push(rcgen::ExtendedKeyUsagePurpose::ClientAuth);
let client_key = rcgen::KeyPair::generate().expect("key generation does not fail");
let client_cert = client_params
.signed_by(&client_key, &ca_cert, &ca_key)
.expect("the client certificate is signed by the CA");
let client_cert_pem = client_cert.pem();
let client_key_pem = client_key.serialize_pem();
let mut roots = rustls::RootCertStore::empty();
roots
.add(ca_cert.der().clone())
.expect("the CA certificate is well-formed DER");
let client_verifier = rustls::server::WebPkiClientVerifier::builder(Arc::new(roots))
.build()
.expect("a verifier with one trusted root builds");
let server_config = rustls::ServerConfig::builder()
.with_client_cert_verifier(client_verifier)
.with_single_cert(vec![server_cert_der], server_key_der)
.expect("the freshly generated server cert and key are valid together");
let server_config = Arc::new(server_config);
let listener = TcpListener::bind("127.0.0.1:0").expect("an ephemeral port is free");
let addr = listener.local_addr().expect("the listener has an address");
std::thread::spawn(move || {
let Ok((mut sock, _)) = listener.accept() else {
return;
};
let Ok(mut conn) = rustls::ServerConnection::new(server_config) else {
return;
};
let mut tls = rustls::Stream::new(&mut conn, &mut sock);
let mut reader = BufReader::new(&mut tls);
let mut request_line = String::new();
if reader.read_line(&mut request_line).unwrap_or(0) == 0 {
return;
}
loop {
let mut header = String::new();
match reader.read_line(&mut header) {
Ok(0) | Err(_) => return,
Ok(_) if header == "\r\n" => break,
Ok(_) => {}
}
}
let _ = tls.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok");
});
(addr, client_cert_pem, client_key_pem)
}
pub(crate) fn start_cookie_server(
routes: Vec<(&'static str, Vec<u8>)>,
) -> (SocketAddr, Arc<std::sync::Mutex<Vec<Option<String>>>>) {
let listener = TcpListener::bind("127.0.0.1:0").expect("an ephemeral port is free");
let addr = listener.local_addr().expect("the listener has an address");
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let stored = seen.clone();
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(stream) = stream else { continue };
let mut writer = stream.try_clone().expect("the socket clones");
let mut reader = BufReader::new(stream);
loop {
let mut request_line = String::new();
match reader.read_line(&mut request_line) {
Ok(0) | Err(_) => return,
Ok(_) => {}
}
let path = request_line
.split_whitespace()
.nth(1)
.unwrap_or("/")
.to_string();
let mut cookie = None;
loop {
let mut header = String::new();
match reader.read_line(&mut header) {
Ok(0) | Err(_) => return,
Ok(_) if header == "\r\n" => break,
Ok(_) => {
if let Some((name, value)) = header.split_once(':') {
if name.trim().eq_ignore_ascii_case("cookie") {
cookie = Some(value.trim().to_string());
}
}
}
}
}
stored.lock().unwrap().push(cookie);
let response = routes
.iter()
.find(|(route, _)| *route == path)
.map(|(_, body)| body.clone())
.unwrap_or_else(|| {
b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n".to_vec()
});
if writer.write_all(&response).is_err() {
return;
}
let _ = writer.flush();
}
}
});
(addr, seen)
}
pub(crate) fn start_proxy_recording_server() -> (SocketAddr, Arc<std::sync::Mutex<Option<String>>>)
{
let listener = TcpListener::bind("127.0.0.1:0").expect("an ephemeral port is free");
let addr = listener.local_addr().expect("the listener has an address");
let seen = Arc::new(std::sync::Mutex::new(None));
let stored = seen.clone();
std::thread::spawn(move || {
let Ok((stream, _)) = listener.accept() else {
return;
};
let mut writer = stream.try_clone().expect("the socket clones");
let mut reader = BufReader::new(stream);
let mut request_line = String::new();
if reader.read_line(&mut request_line).unwrap_or(0) == 0 {
return;
}
loop {
let mut header = String::new();
match reader.read_line(&mut header) {
Ok(0) | Err(_) => return,
Ok(_) if header == "\r\n" => break,
Ok(_) => {}
}
}
*stored.lock().unwrap() = Some(request_line.trim_end().to_string());
let _ = writer.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok");
let _ = writer.flush();
});
(addr, seen)
}