use std::io::{Read, Write};
use super::crypto::{b64_encode, compute_accept};
use super::mask::MaskKeySource;
const MAX_RESPONSE_HEADER_BYTES: usize = 32 * 1024;
const MAX_HEADER_LINE_BYTES: usize = 8 * 1024;
#[derive(Debug, Clone)]
pub(crate) struct Header {
pub name: String,
pub value: String,
}
#[derive(Debug, Clone, Default)]
pub(crate) struct Headers(Vec<Header>);
impl Headers {
pub(crate) fn find_ci(&self, name: &str) -> Option<&str> {
self.0
.iter()
.find(|h| h.name.eq_ignore_ascii_case(name))
.map(|h| h.value.trim())
}
#[cfg(test)]
pub(crate) fn from_pairs<I, S1, S2>(pairs: I) -> Self
where
I: IntoIterator<Item = (S1, S2)>,
S1: Into<String>,
S2: Into<String>,
{
Self(
pairs
.into_iter()
.map(|(name, value)| Header {
name: name.into(),
value: value.into(),
})
.collect(),
)
}
pub(crate) fn header_has_token(&self, name: &str, token: &str) -> bool {
match self.find_ci(name) {
Some(value) => value
.split(',')
.any(|t| t.trim().eq_ignore_ascii_case(token)),
None => false,
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct HttpReject {
pub status: u16,
pub headers: Headers,
#[allow(dead_code)]
pub body: Vec<u8>,
}
#[derive(Debug, Clone)]
pub(crate) struct Handshake {
pub headers: Headers,
pub leftover: Vec<u8>,
}
#[derive(Debug)]
pub(crate) enum HandshakeError {
Io(std::io::Error),
Protocol(String),
HttpStatus(HttpReject),
BadAccept,
}
impl From<std::io::Error> for HandshakeError {
fn from(e: std::io::Error) -> Self {
HandshakeError::Io(e)
}
}
pub(crate) fn upgrade<S: Read + Write>(
stream: &mut S,
host_header: &str,
path: &str,
extra_headers: &[(&'static str, String)],
) -> std::result::Result<Handshake, HandshakeError> {
let key = generate_client_key().map_err(|e| HandshakeError::Protocol(e.0))?;
let expected_accept = compute_accept(&key);
let mut request = Vec::with_capacity(512);
write_request(&mut request, path, host_header, &key, extra_headers);
stream.write_all(&request)?;
stream.flush()?;
let (header_bytes, leftover) = read_response_prefix(stream)?;
let response = parse_response(&header_bytes)
.map_err(|reason| HandshakeError::Protocol(reason.to_string()))?;
if response.status != 101 {
let body = read_response_body(stream, &response.headers, leftover)?;
return Err(HandshakeError::HttpStatus(HttpReject {
status: response.status,
headers: response.headers,
body,
}));
}
if !response.headers.header_has_token("Upgrade", "websocket") {
return Err(HandshakeError::Protocol(
"missing/invalid Upgrade header".into(),
));
}
if !response.headers.header_has_token("Connection", "Upgrade") {
return Err(HandshakeError::Protocol(
"missing/invalid Connection header".into(),
));
}
let accept = response
.headers
.find_ci("Sec-WebSocket-Accept")
.ok_or_else(|| HandshakeError::Protocol("missing Sec-WebSocket-Accept".into()))?;
if accept != expected_accept {
return Err(HandshakeError::BadAccept);
}
Ok(Handshake {
headers: response.headers,
leftover,
})
}
fn generate_client_key() -> Result<String, super::mask::EntropyUnavailable> {
let rng = MaskKeySource::new()?;
let mut bytes = [0u8; 16];
rng.fill(&mut bytes)?;
Ok(b64_encode(&bytes))
}
fn write_request(
out: &mut Vec<u8>,
path: &str,
host_header: &str,
sec_key: &str,
extra_headers: &[(&'static str, String)],
) {
out.extend_from_slice(b"GET ");
out.extend_from_slice(path.as_bytes());
out.extend_from_slice(b" HTTP/1.1\r\n");
push_header(out, "Host", host_header);
push_header(out, "Connection", "Upgrade");
push_header(out, "Upgrade", "websocket");
push_header(out, "Sec-WebSocket-Version", "13");
push_header(out, "Sec-WebSocket-Key", sec_key);
for (name, value) in extra_headers {
push_header(out, name, value);
}
out.extend_from_slice(b"\r\n");
}
fn push_header(out: &mut Vec<u8>, name: &str, value: &str) {
out.extend_from_slice(name.as_bytes());
out.extend_from_slice(b": ");
out.extend_from_slice(value.as_bytes());
out.extend_from_slice(b"\r\n");
}
fn read_response_prefix<S: Read>(
stream: &mut S,
) -> std::result::Result<(Vec<u8>, Vec<u8>), HandshakeError> {
let mut buf = Vec::new();
buf.try_reserve(4096).map_err(|e| {
HandshakeError::Protocol(format!("handshake recv buffer allocation failed: {e}"))
})?;
let mut chunk = [0u8; 4096];
let mut search_from: usize = 0;
loop {
let n = stream.read(&mut chunk)?;
if n == 0 {
return Err(HandshakeError::Protocol(format!(
"server closed during handshake response read (got {} bytes, no `\\r\\n\\r\\n`)",
buf.len()
)));
}
buf.try_reserve(n).map_err(|e| {
HandshakeError::Protocol(format!("handshake recv buffer growth failed: {e}"))
})?;
buf.extend_from_slice(&chunk[..n]);
if buf.len() > MAX_RESPONSE_HEADER_BYTES {
return Err(HandshakeError::Protocol(format!(
"handshake response exceeded {} bytes without `\\r\\n\\r\\n` terminator",
MAX_RESPONSE_HEADER_BYTES
)));
}
let scan_from = search_from.saturating_sub(3);
if let Some(idx) = find_crlf_crlf(&buf[scan_from..]) {
let term_end = scan_from + idx + 4;
let leftover = buf.split_off(term_end);
return Ok((buf, leftover));
}
search_from = buf.len();
}
}
fn find_crlf_crlf(haystack: &[u8]) -> Option<usize> {
haystack.windows(4).position(|w| w == b"\r\n\r\n")
}
#[derive(Debug)]
struct ParsedResponse {
status: u16,
headers: Headers,
}
fn parse_response(bytes: &[u8]) -> std::result::Result<ParsedResponse, &'static str> {
let body_end = bytes
.windows(4)
.position(|w| w == b"\r\n\r\n")
.ok_or("missing \\r\\n\\r\\n terminator")?;
let header_block = &bytes[..body_end];
let mut lines = split_crlf(header_block);
let status_line = lines.next().ok_or("response has no status line")?;
let status = parse_status_line(status_line)?;
let mut headers = Headers::default();
for line in lines {
if line.is_empty() {
continue;
}
if line.len() > MAX_HEADER_LINE_BYTES {
return Err("header line exceeds 8 KiB");
}
if line.starts_with(b" ") || line.starts_with(b"\t") {
return Err("folded header continuation is not supported");
}
let (name, value) = split_header_line(line)?;
headers.0.push(Header { name, value });
}
Ok(ParsedResponse { status, headers })
}
fn split_crlf(bytes: &[u8]) -> impl Iterator<Item = &[u8]> {
bytes.split(|&b| b == b'\n').map(|line| {
if let [body @ .., b'\r'] = line {
body
} else {
line
}
})
}
fn parse_status_line(line: &[u8]) -> std::result::Result<u16, &'static str> {
let s = std::str::from_utf8(line).map_err(|_| "status line is not UTF-8")?;
let mut parts = s.splitn(3, ' ');
let version = parts.next().ok_or("status line missing version")?;
if !version.starts_with("HTTP/1.") {
return Err("status line has non-HTTP/1.x version");
}
let code = parts.next().ok_or("status line missing status code")?;
code.parse::<u16>().map_err(|_| "status code is not a u16")
}
fn split_header_line(line: &[u8]) -> std::result::Result<(String, String), &'static str> {
let colon = line
.iter()
.position(|&b| b == b':')
.ok_or("header line missing `:`")?;
let name = std::str::from_utf8(&line[..colon]).map_err(|_| "header name is not UTF-8")?;
let value = std::str::from_utf8(&line[colon + 1..]).map_err(|_| "header value is not UTF-8")?;
if name.is_empty() || name.chars().any(|c| c.is_ascii_whitespace()) {
return Err("header name has whitespace");
}
Ok((name.to_string(), value.trim().to_string()))
}
fn read_response_body<S: Read>(
stream: &mut S,
headers: &Headers,
leftover: Vec<u8>,
) -> std::result::Result<Vec<u8>, HandshakeError> {
const MAX_BODY_BYTES: usize = 64 * 1024;
let declared_len = headers
.find_ci("Content-Length")
.and_then(|v| v.parse::<usize>().ok());
let Some(content_length) = declared_len else {
return Ok(leftover);
};
if content_length > MAX_BODY_BYTES {
let mut buf = leftover;
let target = MAX_BODY_BYTES.min(content_length);
if buf.len() < target {
let mut tail = try_alloc_zeroed(target - buf.len())?;
let n = read_to_fill(stream, &mut tail)?;
buf.extend_from_slice(&tail[..n]);
}
return Ok(buf);
}
if leftover.len() >= content_length {
return Ok(leftover);
}
let mut buf = leftover;
let want = content_length - buf.len();
let mut tail = try_alloc_zeroed(want)?;
let n = read_to_fill(stream, &mut tail)?;
buf.extend_from_slice(&tail[..n]);
Ok(buf)
}
fn try_alloc_zeroed(n: usize) -> std::result::Result<Vec<u8>, HandshakeError> {
let mut v: Vec<u8> = Vec::new();
v.try_reserve_exact(n).map_err(|e| {
HandshakeError::Protocol(format!("handshake body buffer allocation failed: {e}"))
})?;
v.resize(n, 0u8);
Ok(v)
}
fn read_to_fill<S: Read>(
stream: &mut S,
buf: &mut [u8],
) -> std::result::Result<usize, HandshakeError> {
let mut filled = 0;
while filled < buf.len() {
match stream.read(&mut buf[filled..])? {
0 => break,
n => filled += n,
}
}
Ok(filled)
}
#[cfg(test)]
mod tests {
use super::*;
struct MemStream {
to_read: std::io::Cursor<Vec<u8>>,
written: Vec<u8>,
}
impl MemStream {
fn new(server_bytes: Vec<u8>) -> Self {
Self {
to_read: std::io::Cursor::new(server_bytes),
written: Vec::new(),
}
}
}
impl Read for MemStream {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
self.to_read.read(buf)
}
}
impl Write for MemStream {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.written.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
fn extract_sec_key(req: &[u8]) -> String {
let s = std::str::from_utf8(req).unwrap();
for line in s.split("\r\n") {
if let Some(v) = line.strip_prefix("Sec-WebSocket-Key: ") {
return v.to_string();
}
}
panic!("Sec-WebSocket-Key not in request:\n{s}");
}
#[test]
fn upgrade_signs_with_runtime_key() {
let mut server = MockServer::default();
let result = upgrade(
&mut server,
"host:1234",
"/path",
&[("X-Extra", "abc".into())],
);
assert!(result.is_ok(), "{:?}", result.err());
let req = std::str::from_utf8(&server.written).unwrap();
assert!(req.contains("X-Extra: abc\r\n"), "{req}");
}
#[test]
fn rejects_when_accept_mismatch() {
let resp = b"\
HTTP/1.1 101 Switching Protocols\r\n\
Upgrade: websocket\r\n\
Connection: Upgrade\r\n\
Sec-WebSocket-Accept: bogus=\r\n\r\n";
let mut server = MemStream::new(resp.to_vec());
let err = upgrade(&mut server, "host:1", "/", &[]).unwrap_err();
assert!(matches!(err, HandshakeError::BadAccept), "{err:?}");
}
#[test]
fn surfaces_4xx_as_http_status() {
let resp = b"\
HTTP/1.1 401 Unauthorized\r\n\
WWW-Authenticate: Basic\r\n\
Content-Length: 11\r\n\r\nhello world";
let mut server = MemStream::new(resp.to_vec());
let err = upgrade(&mut server, "host:1", "/", &[]).unwrap_err();
match err {
HandshakeError::HttpStatus(reject) => {
assert_eq!(reject.status, 401);
assert_eq!(reject.body, b"hello world");
assert_eq!(reject.headers.find_ci("WWW-Authenticate"), Some("Basic"));
}
other => panic!("expected HttpStatus, got {other:?}"),
}
}
#[test]
fn rejects_when_missing_upgrade_header() {
let mut server = MockServer::without_header("Upgrade");
let err = upgrade(&mut server, "host:1", "/", &[]).unwrap_err();
assert!(matches!(err, HandshakeError::Protocol(_)), "{err:?}");
}
#[test]
fn rejects_when_missing_connection_header() {
let mut server = MockServer::without_header("Connection");
let err = upgrade(&mut server, "host:1", "/", &[]).unwrap_err();
assert!(matches!(err, HandshakeError::Protocol(_)), "{err:?}");
}
#[test]
fn parse_status_line_minimal() {
assert_eq!(
parse_status_line(b"HTTP/1.1 101 Switching Protocols").unwrap(),
101
);
assert_eq!(parse_status_line(b"HTTP/1.0 200 OK").unwrap(), 200);
assert_eq!(
parse_status_line(b"HTTP/1.1 421 Misdirected Request").unwrap(),
421
);
}
#[test]
fn parse_status_line_rejects_garbage() {
assert!(parse_status_line(b"GARBAGE").is_err());
assert!(parse_status_line(b"HTTP/2.0 200 OK").is_err());
assert!(parse_status_line(b"HTTP/1.1 abc OK").is_err());
}
#[test]
fn slow_loris_cap() {
let garbage = vec![b'A'; 33 * 1024];
let mut server = MemStream::new(garbage);
let err = upgrade(&mut server, "host:1", "/", &[]).unwrap_err();
assert!(
matches!(&err, HandshakeError::Protocol(m) if m.contains("exceeded")),
"{err:?}"
);
}
#[test]
fn terminator_straddles_read_boundary() {
let mut buf = b"HTTP/1.1 101 OK\r\nA: 1\r".to_vec();
assert!(find_crlf_crlf(&buf).is_none());
buf.extend_from_slice(b"\n\r\n");
let idx = find_crlf_crlf(&buf).expect("must find terminator across boundary");
assert_eq!(idx, buf.len() - 4);
}
struct MockServer {
written: Vec<u8>,
to_send: std::io::Cursor<Vec<u8>>,
prepared: bool,
omit_header: Option<&'static str>,
}
impl MockServer {
fn without_header(name: &'static str) -> Self {
Self {
written: Vec::new(),
to_send: std::io::Cursor::new(Vec::new()),
prepared: false,
omit_header: Some(name),
}
}
fn prepare_response(&mut self) {
let key = extract_sec_key(&self.written);
let omit = self.omit_header;
let extras = &[];
let resp = build_response(&key, omit, extras);
self.to_send = std::io::Cursor::new(resp);
self.prepared = true;
}
}
impl Default for MockServer {
fn default() -> Self {
Self {
written: Vec::new(),
to_send: std::io::Cursor::new(Vec::new()),
prepared: false,
omit_header: None,
}
}
}
fn build_response(client_key: &str, omit: Option<&str>, extras: &[(&str, &str)]) -> Vec<u8> {
let accept = compute_accept(client_key);
let mut resp = String::new();
resp.push_str("HTTP/1.1 101 Switching Protocols\r\n");
if omit != Some("Upgrade") {
resp.push_str("Upgrade: websocket\r\n");
}
if omit != Some("Connection") {
resp.push_str("Connection: Upgrade\r\n");
}
resp.push_str(&format!("Sec-WebSocket-Accept: {accept}\r\n"));
for (k, v) in extras {
resp.push_str(&format!("{k}: {v}\r\n"));
}
resp.push_str("\r\n");
resp.into_bytes()
}
impl Read for MockServer {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if !self.prepared {
self.prepare_response();
}
self.to_send.read(buf)
}
}
impl Write for MockServer {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.written.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
}