diff --git a/crates/net/src/client.rs b/crates/net/src/client.rs index d80b3fa..49f5275 100644 --- a/crates/net/src/client.rs +++ b/crates/net/src/client.rs @@ -1,7 +1,9 @@ -//! High-level HTTP/1.1 client with connection pooling. +//! High-level HTTP client with connection pooling and HTTP/2 support. //! -//! Brings together TCP, TLS 1.3, DNS, URL parsing, and HTTP message -//! parsing into a single `HttpClient` that can fetch HTTP and HTTPS URLs. +//! Brings together TCP, TLS 1.3, DNS, URL parsing, HTTP/1.1, and HTTP/2 +//! into a single `HttpClient` that can fetch HTTP and HTTPS URLs. +//! When connecting over TLS, ALPN negotiation automatically selects HTTP/2 +//! if the server supports it, falling back to HTTP/1.1 otherwise. use std::collections::HashMap; use std::fmt; @@ -12,6 +14,8 @@ use we_url::Url; use crate::cookie::{CookieJar, RequestContext}; use crate::http::{self, Headers, HttpResponse, Method}; +use crate::http2::connection::Http2Connection; +use crate::http2::frame::Http2Error; use crate::tcp::{self, TcpConnection}; use crate::tls::handshake::{self, HandshakeError, TlsStream}; @@ -45,6 +49,8 @@ pub enum ClientError { Http(http::HttpError), /// Too many redirects. TooManyRedirects, + /// HTTP/2 protocol error. + Http2(Http2Error), /// Connection was closed unexpectedly. ConnectionClosed, /// I/O error. @@ -60,6 +66,7 @@ impl fmt::Display for ClientError { Self::Tls(e) => write!(f, "TLS error: {e}"), Self::Http(e) => write!(f, "HTTP error: {e}"), Self::TooManyRedirects => write!(f, "too many redirects"), + Self::Http2(e) => write!(f, "HTTP/2 error: {e}"), Self::ConnectionClosed => write!(f, "connection closed"), Self::Io(e) => write!(f, "I/O error: {e}"), } @@ -90,6 +97,12 @@ impl From for ClientError { } } +impl From for ClientError { + fn from(e: Http2Error) -> Self { + Self::Http2(e) + } +} + pub type Result = std::result::Result; // --------------------------------------------------------------------------- @@ -201,9 +214,17 @@ impl ConnectionPool { // HttpClient // --------------------------------------------------------------------------- -/// High-level HTTP/1.1 client with connection pooling, redirect following, and cookie jar. +/// Key for HTTP/2 connection pooling (one multiplexed connection per origin). +#[derive(Hash, Eq, PartialEq, Clone, Debug)] +struct H2ConnectionKey { + host: String, + port: u16, +} + +/// High-level HTTP client with connection pooling, HTTP/2 support, redirect following, and cookie jar. pub struct HttpClient { pool: ConnectionPool, + h2_connections: HashMap>>, max_redirects: u32, connect_timeout: Duration, read_timeout: Duration, @@ -215,6 +236,7 @@ impl HttpClient { pub fn new() -> Self { Self { pool: ConnectionPool::new(DEFAULT_MAX_IDLE_TIME, DEFAULT_MAX_PER_HOST), + h2_connections: HashMap::new(), max_redirects: DEFAULT_MAX_REDIRECTS, connect_timeout: DEFAULT_CONNECT_TIMEOUT, read_timeout: DEFAULT_READ_TIMEOUT, @@ -338,29 +360,99 @@ impl HttpClient { } } + // Try HTTP/2 for TLS connections + if is_tls { + let h2_key = H2ConnectionKey { + host: host.clone(), + port, + }; + + // Check if we have an existing HTTP/2 connection for this origin + if self.h2_connections.contains_key(&h2_key) { + let response = self.execute_h2_request( + &h2_key, + method, + &path, + &host, + &merged_headers, + body, + url, + ); + + match response { + Ok(resp) => return Ok(resp), + Err(ClientError::Http2(_)) => { + // HTTP/2 connection failed — remove it and fall through to new connection + self.h2_connections.remove(&h2_key); + } + Err(e) => return Err(e), + } + } + + // Try to establish a new connection with ALPN + let tcp = TcpConnection::connect_timeout(&host, port, self.connect_timeout)?; + let (tls, alpn) = handshake::connect_with_alpn(tcp, &host, &["h2", "http/1.1"])?; + + if alpn.as_deref() == Some("h2") { + // HTTP/2 negotiated — create HTTP/2 connection + let h2_conn = Http2Connection::new(tls)?; + self.h2_connections.insert(h2_key.clone(), h2_conn); + + return self.execute_h2_request( + &h2_key, + method, + &path, + &host, + &merged_headers, + body, + url, + ); + } + + // HTTP/1.1 over TLS — use normal path + let conn = Connection::Tls(tls); + return self.execute_h1_request(conn, method, &path, &host, &merged_headers, body, url); + } + + // Plain HTTP (no TLS) — always HTTP/1.1 let key = ConnectionKey { host: host.clone(), port, - is_tls, + is_tls: false, }; - // Try to reuse a pooled connection, fall back to new connection - let mut conn = match self.pool.take(&key) { + let conn = match self.pool.take(&key) { Some(conn) => conn, - None => self.connect(&host, port, is_tls)?, + None => { + let tcp = TcpConnection::connect_timeout(&host, port, self.connect_timeout)?; + Connection::Plain(tcp) + } }; + self.execute_h1_request(conn, method, &path, &host, &merged_headers, body, url) + } + + /// Execute an HTTP/1.1 request over the given connection. + #[allow(clippy::too_many_arguments)] + fn execute_h1_request( + &mut self, + mut conn: Connection, + method: Method, + path: &str, + host: &str, + headers: &Headers, + body: Option<&[u8]>, + url: &Url, + ) -> Result { conn.set_read_timeout(Some(self.read_timeout))?; - // Serialize and send request - let request_bytes = http::serialize_request(method, &path, &host, &merged_headers, body); + let request_bytes = http::serialize_request(method, path, host, headers, body); conn.write_all(&request_bytes)?; conn.flush()?; - // Read and parse response let response = read_response(&mut conn)?; - // Store Set-Cookie headers from the response. + // Store Set-Cookie headers let set_cookies: Vec = response .headers .get_all("Set-Cookie") @@ -373,22 +465,78 @@ impl HttpClient { // Return connection to pool if keep-alive if !response.connection_close() { + // Determine the connection key for pooling + let is_tls = matches!(conn, Connection::Tls(_)); + let port = url + .port_or_default() + .unwrap_or(if is_tls { 443 } else { 80 }); + let key = ConnectionKey { + host: host.to_string(), + port, + is_tls, + }; self.pool.put(key, conn); } Ok(response) } - /// Establish a new connection (plain TCP or TLS). - fn connect(&self, host: &str, port: u16, is_tls: bool) -> Result { - let tcp = TcpConnection::connect_timeout(host, port, self.connect_timeout)?; + /// Execute a request over an existing HTTP/2 connection. + #[allow(clippy::too_many_arguments)] + fn execute_h2_request( + &mut self, + h2_key: &H2ConnectionKey, + method: Method, + path: &str, + authority: &str, + headers: &Headers, + body: Option<&[u8]>, + url: &Url, + ) -> Result { + let extra_headers: Vec<(String, String)> = headers + .iter() + .filter(|(name, _)| { + // Skip pseudo-headers and host (authority is used instead) + let lower = name.to_ascii_lowercase(); + lower != "host" && lower != "connection" && lower != "transfer-encoding" + }) + .map(|(name, value)| (name.to_ascii_lowercase(), value.to_string())) + .collect(); - if is_tls { - let tls = handshake::connect(tcp, host)?; - Ok(Connection::Tls(tls)) - } else { - Ok(Connection::Plain(tcp)) + let h2_conn = self.h2_connections.get_mut(h2_key).unwrap(); + + let stream_id = + h2_conn.send_request(method.as_str(), path, authority, &extra_headers, body)?; + + let (resp_headers, resp_body, status_code) = h2_conn.read_response(stream_id)?; + + // Convert HTTP/2 response to HttpResponse + let mut response_headers = Headers::new(); + for (name, value) in &resp_headers { + let name_str = String::from_utf8_lossy(name); + let value_str = String::from_utf8_lossy(value); + if !name_str.starts_with(':') { + response_headers.add(&name_str, &value_str); + } } + + // Store Set-Cookie headers + let set_cookies: Vec = response_headers + .get_all("Set-Cookie") + .into_iter() + .map(|s| s.to_string()) + .collect(); + for header in &set_cookies { + self.cookie_jar.store_from_header(header, url); + } + + Ok(HttpResponse { + version: "HTTP/2".to_string(), + status_code, + reason: reason_phrase(status_code).to_string(), + headers: response_headers, + body: resp_body, + }) } } @@ -553,6 +701,29 @@ fn determine_body_strategy(headers: &str, status_code: u16) -> BodyStrategy { BodyStrategy::ReadUntilClose } +/// Standard HTTP reason phrase for a status code. +fn reason_phrase(status: u16) -> &'static str { + match status { + 200 => "OK", + 201 => "Created", + 204 => "No Content", + 301 => "Moved Permanently", + 302 => "Found", + 304 => "Not Modified", + 307 => "Temporary Redirect", + 308 => "Permanent Redirect", + 400 => "Bad Request", + 401 => "Unauthorized", + 403 => "Forbidden", + 404 => "Not Found", + 405 => "Method Not Allowed", + 500 => "Internal Server Error", + 502 => "Bad Gateway", + 503 => "Service Unavailable", + _ => "", + } +} + /// Check if chunked body data contains the terminating `0\r\n\r\n`. fn has_chunked_terminator(data: &[u8]) -> bool { // Look for \r\n0\r\n\r\n (the final chunk after some data) or 0\r\n\r\n at start diff --git a/crates/net/src/http2/connection.rs b/crates/net/src/http2/connection.rs new file mode 100644 index 0000000..bc998de --- /dev/null +++ b/crates/net/src/http2/connection.rs @@ -0,0 +1,898 @@ +//! HTTP/2 connection management (RFC 7540). +//! +//! Manages the HTTP/2 connection lifecycle: preface exchange, SETTINGS +//! negotiation, stream multiplexing, flow control, and request/response. + +use std::collections::HashMap; +use std::io::{Read, Write}; + +use super::frame::{ + self, data_frame, goaway_frame, headers_frame, ping_frame, settings_ack_frame, settings_frame, + window_update_frame, ErrorCode, Frame, Http2Error, Result, CONNECTION_PREFACE, + DEFAULT_MAX_FRAME_SIZE, FLAG_ACK, FLAG_END_HEADERS, FLAG_END_STREAM, FRAME_CONTINUATION, + FRAME_DATA, FRAME_GOAWAY, FRAME_HEADERS, FRAME_PING, FRAME_RST_STREAM, FRAME_SETTINGS, + FRAME_WINDOW_UPDATE, SETTINGS_ENABLE_PUSH, SETTINGS_HEADER_TABLE_SIZE, + SETTINGS_INITIAL_WINDOW_SIZE, SETTINGS_MAX_CONCURRENT_STREAMS, SETTINGS_MAX_FRAME_SIZE, + SETTINGS_MAX_HEADER_LIST_SIZE, +}; +use super::hpack::{Decoder, Encoder, HeaderField}; +use super::stream::{Stream, StreamState, DEFAULT_INITIAL_WINDOW_SIZE}; + +/// HTTP/2 response: (headers as (name, value) pairs, body bytes, status code). +pub type H2Response = (Vec<(Vec, Vec)>, Vec, u16); + +// --------------------------------------------------------------------------- +// Connection-level settings +// --------------------------------------------------------------------------- + +/// Peer's settings (received from the server). +struct PeerSettings { + header_table_size: u32, + enable_push: bool, + max_concurrent_streams: u32, + initial_window_size: u32, + max_frame_size: u32, + max_header_list_size: u32, +} + +impl Default for PeerSettings { + fn default() -> Self { + Self { + header_table_size: 4096, + enable_push: true, + max_concurrent_streams: u32::MAX, + initial_window_size: DEFAULT_INITIAL_WINDOW_SIZE, + max_frame_size: DEFAULT_MAX_FRAME_SIZE, + max_header_list_size: u32::MAX, + } + } +} + +// --------------------------------------------------------------------------- +// HTTP/2 Connection +// --------------------------------------------------------------------------- + +/// An HTTP/2 connection over a TLS stream. +/// +/// Manages stream multiplexing, flow control, HPACK encoding/decoding, +/// and the SETTINGS handshake. +pub struct Http2Connection { + /// Underlying transport (TLS stream). + stream: S, + /// Active streams indexed by stream ID. + streams: HashMap, + /// Next client-initiated stream ID (odd, starting at 1). + next_stream_id: u32, + /// Connection-level send window. + conn_send_window: i64, + /// Connection-level receive window. + conn_recv_window: i64, + /// Peer's (server's) settings. + peer_settings: PeerSettings, + /// HPACK encoder for request headers. + hpack_encoder: Encoder, + /// HPACK decoder for response headers. + hpack_decoder: Decoder, + /// Pending encoder table size update (from peer SETTINGS). + pending_table_size_update: Option, + /// Whether we have received the initial SETTINGS from the peer. + settings_received: bool, + /// Whether a GOAWAY has been received. + goaway_received: bool, + /// Last stream ID from GOAWAY. + goaway_last_stream_id: u32, +} + +impl Http2Connection { + /// Perform the HTTP/2 connection handshake. + /// + /// Sends the connection preface and initial SETTINGS, then reads the + /// server's SETTINGS and acknowledges it. + pub fn new(mut stream: S) -> Result { + // Send connection preface + stream + .write_all(CONNECTION_PREFACE) + .map_err(Http2Error::Io)?; + + // Send our SETTINGS + let our_settings = [ + (SETTINGS_MAX_CONCURRENT_STREAMS, 100), + (SETTINGS_INITIAL_WINDOW_SIZE, DEFAULT_INITIAL_WINDOW_SIZE), + (SETTINGS_ENABLE_PUSH, 0), // Disable server push + ]; + let settings = settings_frame(&our_settings, false); + settings.write_to(&mut stream)?; + + let mut conn = Self { + stream, + streams: HashMap::new(), + next_stream_id: 1, + conn_send_window: DEFAULT_INITIAL_WINDOW_SIZE as i64, + conn_recv_window: DEFAULT_INITIAL_WINDOW_SIZE as i64, + peer_settings: PeerSettings::default(), + hpack_encoder: Encoder::new(4096), + hpack_decoder: Decoder::new(4096), + pending_table_size_update: None, + settings_received: false, + goaway_received: false, + goaway_last_stream_id: 0, + }; + + // Read frames until we get the server's SETTINGS + conn.read_until_settings()?; + + // Send a larger connection-level window update to allow more data + // Default is 65535; bump to ~1MB for better throughput + let window_bump = 1_048_576 - DEFAULT_INITIAL_WINDOW_SIZE; + if window_bump > 0 { + let wu = window_update_frame(0, window_bump); + wu.write_to(&mut conn.stream)?; + conn.conn_recv_window += window_bump as i64; + } + + Ok(conn) + } + + /// Send an HTTP request and return the stream ID. + /// + /// The request headers are HPACK-encoded and sent as a HEADERS frame. + /// If a body is provided, it is sent as DATA frame(s). + pub fn send_request( + &mut self, + method: &str, + path: &str, + authority: &str, + extra_headers: &[(String, String)], + body: Option<&[u8]>, + ) -> Result { + if self.goaway_received { + return Err(Http2Error::Protocol("connection received GOAWAY".into())); + } + + let stream_id = self.next_stream_id; + self.next_stream_id += 2; + + // Build pseudo-headers + regular headers + let mut headers = vec![ + HeaderField::new(b":method", method.as_bytes()), + HeaderField::new(b":scheme", b"https"), + HeaderField::new(b":path", path.as_bytes()), + HeaderField::new(b":authority", authority.as_bytes()), + ]; + for (name, value) in extra_headers { + headers.push(HeaderField::new(name.as_bytes(), value.as_bytes())); + } + + // Encode headers with HPACK + let mut header_block = Vec::new(); + if let Some(new_size) = self.pending_table_size_update.take() { + self.hpack_encoder + .encode_table_size_update(&mut header_block, new_size); + } + header_block.extend_from_slice(&self.hpack_encoder.encode(&headers)); + + let end_stream = body.is_none(); + + // Create stream and transition to open + let mut stream = Stream::new( + stream_id, + self.peer_settings.initial_window_size, + DEFAULT_INITIAL_WINDOW_SIZE, + ); + stream.send_headers()?; + if end_stream { + stream.send_end_stream()?; + } + self.streams.insert(stream_id, stream); + + // Send HEADERS frame + let hdr_frame = headers_frame(stream_id, header_block, end_stream); + hdr_frame.write_to(&mut self.stream)?; + + // Send body as DATA frames if present + if let Some(data) = body { + self.send_data(stream_id, data)?; + } + + Ok(stream_id) + } + + /// Read the response for a given stream. + /// + /// Reads frames until the stream is complete (END_STREAM received on + /// both headers and data). Returns the response headers and body. + pub fn read_response(&mut self, stream_id: u32) -> Result { + // Read frames until this stream is done + loop { + let stream = self + .streams + .get(&stream_id) + .ok_or_else(|| Http2Error::Protocol(format!("stream {stream_id} not found")))?; + + if stream.is_closed() || stream.state == StreamState::HalfClosedRemote { + break; + } + + self.read_and_process_frame()?; + } + + let stream = self + .streams + .get(&stream_id) + .ok_or_else(|| Http2Error::Protocol(format!("stream {stream_id} not found")))?; + + let status = stream.status_code.unwrap_or(0); + Ok((stream.response_headers.clone(), stream.body.clone(), status)) + } + + /// Close the connection gracefully by sending GOAWAY. + pub fn close(&mut self) -> Result<()> { + let last_id = if self.next_stream_id > 1 { + self.next_stream_id - 2 + } else { + 0 + }; + let frame = goaway_frame(last_id, ErrorCode::NoError); + frame.write_to(&mut self.stream)?; + Ok(()) + } + + // ----------------------------------------------------------------------- + // Internal: frame sending + // ----------------------------------------------------------------------- + + /// Send body data as DATA frame(s) with flow control. + fn send_data(&mut self, stream_id: u32, data: &[u8]) -> Result<()> { + let mut offset = 0; + while offset < data.len() { + let remaining = data.len() - offset; + // Respect both connection and stream flow control windows + let stream = self + .streams + .get(&stream_id) + .ok_or_else(|| Http2Error::Protocol("stream not found".into()))?; + + let max_by_stream = stream.send_window.max(0) as usize; + let max_by_conn = self.conn_send_window.max(0) as usize; + let max_by_frame = self.peer_settings.max_frame_size as usize; + let chunk_size = remaining + .min(max_by_stream) + .min(max_by_conn) + .min(max_by_frame); + + if chunk_size == 0 { + // Flow control window exhausted — read frames to get WINDOW_UPDATE + self.read_and_process_frame()?; + continue; + } + + let is_last = offset + chunk_size >= data.len(); + let chunk = &data[offset..offset + chunk_size]; + let frame = data_frame(stream_id, chunk.to_vec(), is_last); + frame.write_to(&mut self.stream)?; + + // Update flow control + self.conn_send_window -= chunk_size as i64; + let stream = self.streams.get_mut(&stream_id).unwrap(); + stream.consume_send_window(chunk_size as u32)?; + if is_last { + stream.send_end_stream()?; + } + + offset += chunk_size; + } + Ok(()) + } + + // ----------------------------------------------------------------------- + // Internal: frame reading + // ----------------------------------------------------------------------- + + /// Read frames until we get the initial SETTINGS from the server. + /// Also consumes the SETTINGS ACK for our settings if present. + fn read_until_settings(&mut self) -> Result<()> { + let mut got_settings = false; + let mut got_settings_ack = false; + + for _ in 0..32 { + if got_settings && got_settings_ack { + return Ok(()); + } + let frame = Frame::read_from(&mut self.stream, self.peer_settings.max_frame_size)?; + + if frame.header.frame_type == FRAME_SETTINGS { + if frame.header.has_flag(FLAG_ACK) { + got_settings_ack = true; + // Process it (no-op for ACK, but keeps the flow clean) + self.process_frame(frame)?; + } else { + got_settings = true; + self.process_frame(frame)?; + } + } else { + // Process non-SETTINGS frames normally + self.process_frame(frame)?; + } + } + + if !self.settings_received { + return Err(Http2Error::Protocol( + "did not receive SETTINGS from server".into(), + )); + } + Ok(()) + } + + /// Read one frame from the transport and process it. + fn read_and_process_frame(&mut self) -> Result<()> { + let frame = Frame::read_from(&mut self.stream, self.peer_settings.max_frame_size)?; + self.process_frame(frame) + } + + /// Process a received frame. + fn process_frame(&mut self, frame: Frame) -> Result<()> { + match frame.header.frame_type { + FRAME_SETTINGS => self.handle_settings(&frame), + FRAME_HEADERS => self.handle_headers(&frame), + FRAME_CONTINUATION => self.handle_continuation(&frame), + FRAME_DATA => self.handle_data(&frame), + FRAME_WINDOW_UPDATE => self.handle_window_update(&frame), + FRAME_RST_STREAM => self.handle_rst_stream(&frame), + FRAME_GOAWAY => self.handle_goaway(&frame), + FRAME_PING => self.handle_ping(&frame), + _ => { + // Unknown frame types MUST be ignored (RFC 7540 §4.1) + Ok(()) + } + } + } + + fn handle_settings(&mut self, frame: &Frame) -> Result<()> { + if frame.header.stream_id != 0 { + return Err(Http2Error::Protocol("SETTINGS on non-zero stream".into())); + } + + if frame.header.has_flag(FLAG_ACK) { + // ACK for our SETTINGS — nothing more to do + return Ok(()); + } + + let settings = frame::parse_settings(&frame.payload)?; + for s in &settings { + match s.id { + SETTINGS_HEADER_TABLE_SIZE => { + self.peer_settings.header_table_size = s.value; + self.pending_table_size_update = Some(s.value as usize); + } + SETTINGS_ENABLE_PUSH => { + self.peer_settings.enable_push = s.value != 0; + } + SETTINGS_MAX_CONCURRENT_STREAMS => { + self.peer_settings.max_concurrent_streams = s.value; + } + SETTINGS_INITIAL_WINDOW_SIZE => { + if s.value > 0x7FFF_FFFF { + return Err(Http2Error::FlowControl); + } + // Adjust existing streams' send windows + let delta = s.value as i64 - self.peer_settings.initial_window_size as i64; + for stream in self.streams.values_mut() { + stream.send_window += delta; + } + self.peer_settings.initial_window_size = s.value; + } + SETTINGS_MAX_FRAME_SIZE => { + if s.value < DEFAULT_MAX_FRAME_SIZE || s.value > frame::MAX_FRAME_SIZE_LIMIT { + return Err(Http2Error::Protocol( + "invalid SETTINGS_MAX_FRAME_SIZE".into(), + )); + } + self.peer_settings.max_frame_size = s.value; + } + SETTINGS_MAX_HEADER_LIST_SIZE => { + self.peer_settings.max_header_list_size = s.value; + } + _ => {} // Unknown settings MUST be ignored (RFC 7540 §6.5.2) + } + } + + self.settings_received = true; + + // Send SETTINGS ACK + let ack = settings_ack_frame(); + ack.write_to(&mut self.stream)?; + + Ok(()) + } + + fn handle_headers(&mut self, frame: &Frame) -> Result<()> { + let stream_id = frame.header.stream_id; + if stream_id == 0 { + return Err(Http2Error::Protocol("HEADERS on stream 0".into())); + } + + let stream = match self.streams.get_mut(&stream_id) { + Some(s) => s, + None => { + // Server-initiated stream (even ID) or unknown — ignore + return Ok(()); + } + }; + + stream.header_block.extend_from_slice(&frame.payload); + + if frame.header.has_flag(FLAG_END_HEADERS) { + self.decode_headers(stream_id)?; + } + + if frame.header.has_flag(FLAG_END_STREAM) { + let stream = self.streams.get_mut(&stream_id).unwrap(); + stream.recv_end_stream()?; + } + + Ok(()) + } + + fn handle_continuation(&mut self, frame: &Frame) -> Result<()> { + let stream_id = frame.header.stream_id; + let stream = match self.streams.get_mut(&stream_id) { + Some(s) => s, + None => return Ok(()), + }; + + stream.header_block.extend_from_slice(&frame.payload); + + if frame.header.has_flag(FLAG_END_HEADERS) { + self.decode_headers(stream_id)?; + } + + Ok(()) + } + + /// Decode accumulated header block for a stream. + fn decode_headers(&mut self, stream_id: u32) -> Result<()> { + let stream = self.streams.get_mut(&stream_id).unwrap(); + let header_block = std::mem::take(&mut stream.header_block); + + let headers = self.hpack_decoder.decode(&header_block)?; + + for hf in &headers { + if hf.name == b":status" { + if let Ok(s) = std::str::from_utf8(&hf.value) { + stream.status_code = s.parse().ok(); + } + } + } + + stream.response_headers = headers.into_iter().map(|hf| (hf.name, hf.value)).collect(); + + Ok(()) + } + + fn handle_data(&mut self, frame: &Frame) -> Result<()> { + let stream_id = frame.header.stream_id; + if stream_id == 0 { + return Err(Http2Error::Protocol("DATA on stream 0".into())); + } + + let data_len = frame.payload.len() as u32; + + // Connection-level flow control + self.conn_recv_window -= data_len as i64; + if self.conn_recv_window < 0 { + return Err(Http2Error::FlowControl); + } + + let stream = match self.streams.get_mut(&stream_id) { + Some(s) => s, + None => { + // Stream may have been reset or doesn't exist — send connection-level + // WINDOW_UPDATE to reclaim the window, then ignore. + let wu = window_update_frame(0, data_len); + wu.write_to(&mut self.stream)?; + self.conn_recv_window += data_len as i64; + return Ok(()); + } + }; + + stream.consume_recv_window(data_len)?; + stream.body.extend_from_slice(&frame.payload); + + if frame.header.has_flag(FLAG_END_STREAM) { + stream.recv_end_stream()?; + } + + // Send WINDOW_UPDATE for stream and connection to keep data flowing + if data_len > 0 { + let stream_wu = window_update_frame(stream_id, data_len); + stream_wu.write_to(&mut self.stream)?; + let stream = self.streams.get_mut(&stream_id).unwrap(); + stream.increase_recv_window(data_len); + + let conn_wu = window_update_frame(0, data_len); + conn_wu.write_to(&mut self.stream)?; + self.conn_recv_window += data_len as i64; + } + + Ok(()) + } + + fn handle_window_update(&mut self, frame: &Frame) -> Result<()> { + if frame.payload.len() != 4 { + return Err(Http2Error::Protocol( + "WINDOW_UPDATE payload must be 4 bytes".into(), + )); + } + let increment = u32::from_be_bytes([ + frame.payload[0], + frame.payload[1], + frame.payload[2], + frame.payload[3], + ]) & 0x7FFF_FFFF; + + if increment == 0 { + return Err(Http2Error::Protocol( + "WINDOW_UPDATE increment must be non-zero".into(), + )); + } + + if frame.header.stream_id == 0 { + self.conn_send_window += increment as i64; + if self.conn_send_window > 0x7FFF_FFFF { + return Err(Http2Error::FlowControl); + } + } else if let Some(stream) = self.streams.get_mut(&frame.header.stream_id) { + stream.increase_send_window(increment)?; + } + + Ok(()) + } + + fn handle_rst_stream(&mut self, frame: &Frame) -> Result<()> { + if frame.payload.len() != 4 { + return Err(Http2Error::Protocol( + "RST_STREAM payload must be 4 bytes".into(), + )); + } + let error_code = u32::from_be_bytes([ + frame.payload[0], + frame.payload[1], + frame.payload[2], + frame.payload[3], + ]); + let error_code = ErrorCode::from_u32(error_code); + + if let Some(stream) = self.streams.get_mut(&frame.header.stream_id) { + // Don't propagate the error — just mark the stream as closed. + // The caller will see the error when reading the response. + stream.state = StreamState::Closed; + // Store the error for later retrieval, but don't fail the whole connection. + let _ = stream.recv_rst_stream(error_code); + } + + Ok(()) + } + + fn handle_goaway(&mut self, frame: &Frame) -> Result<()> { + if frame.payload.len() < 8 { + return Err(Http2Error::Protocol("GOAWAY payload too short".into())); + } + let last_stream_id = u32::from_be_bytes([ + frame.payload[0], + frame.payload[1], + frame.payload[2], + frame.payload[3], + ]) & 0x7FFF_FFFF; + let error_code = u32::from_be_bytes([ + frame.payload[4], + frame.payload[5], + frame.payload[6], + frame.payload[7], + ]); + let error_code = ErrorCode::from_u32(error_code); + + self.goaway_received = true; + self.goaway_last_stream_id = last_stream_id; + + // Close streams with ID > last_stream_id + for (id, stream) in &mut self.streams { + if *id > last_stream_id { + stream.state = StreamState::Closed; + } + } + + if error_code != ErrorCode::NoError { + let debug_data = if frame.payload.len() > 8 { + String::from_utf8_lossy(&frame.payload[8..]).to_string() + } else { + String::new() + }; + return Err(Http2Error::GoAway(error_code, debug_data)); + } + + Ok(()) + } + + fn handle_ping(&mut self, frame: &Frame) -> Result<()> { + if frame.header.stream_id != 0 { + return Err(Http2Error::Protocol("PING on non-zero stream".into())); + } + if frame.payload.len() != 8 { + return Err(Http2Error::Protocol("PING payload must be 8 bytes".into())); + } + + if frame.header.has_flag(FLAG_ACK) { + // PING ACK — we don't track outgoing pings, so just ignore + return Ok(()); + } + + // Respond with PING ACK containing same payload + let mut data = [0u8; 8]; + data.copy_from_slice(&frame.payload); + let ack = ping_frame(data, true); + ack.write_to(&mut self.stream)?; + + Ok(()) + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::super::frame::FrameHeader; + use super::*; + + /// A mock transport that records writes and replays reads. + struct MockTransport { + /// Data to be read by the connection. + read_buf: Vec, + read_pos: usize, + /// Data written by the connection. + write_buf: Vec, + } + + impl MockTransport { + fn new(read_data: Vec) -> Self { + Self { + read_buf: read_data, + read_pos: 0, + write_buf: Vec::new(), + } + } + + fn written(&self) -> &[u8] { + &self.write_buf + } + } + + impl Read for MockTransport { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + if self.read_pos >= self.read_buf.len() { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "no more data")); + } + let available = &self.read_buf[self.read_pos..]; + let to_copy = available.len().min(buf.len()); + buf[..to_copy].copy_from_slice(&available[..to_copy]); + self.read_pos += to_copy; + Ok(to_copy) + } + } + + impl Write for MockTransport { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.write_buf.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + + use std::io; + + /// Build server preface: SETTINGS frame + SETTINGS ACK + fn build_server_preface(settings: &[(u16, u32)]) -> Vec { + let mut buf = Vec::new(); + let sf = settings_frame(settings, false); + sf.write_to(&mut buf).unwrap(); + // Also write a SETTINGS ACK for our SETTINGS + let ack = settings_ack_frame(); + ack.write_to(&mut buf).unwrap(); + buf + } + + #[test] + fn connection_handshake() { + let server_data = build_server_preface(&[ + (SETTINGS_MAX_CONCURRENT_STREAMS, 128), + (SETTINGS_INITIAL_WINDOW_SIZE, 65535), + ]); + + let transport = MockTransport::new(server_data); + let conn = Http2Connection::new(transport).unwrap(); + + assert!(conn.settings_received); + assert_eq!(conn.peer_settings.max_concurrent_streams, 128); + assert_eq!(conn.next_stream_id, 1); + } + + #[test] + fn connection_sends_preface() { + let server_data = build_server_preface(&[]); + let transport = MockTransport::new(server_data); + let conn = Http2Connection::new(transport).unwrap(); + + // Verify connection preface was sent + let written = &conn.stream.write_buf; + assert!(written.starts_with(CONNECTION_PREFACE)); + } + + #[test] + fn send_request_allocates_stream_id() { + let server_data = build_server_preface(&[]); + let transport = MockTransport::new(server_data); + let mut conn = Http2Connection::new(transport).unwrap(); + + // We can't fully test send_request without a response, but we can + // verify stream ID allocation by checking next_stream_id advances. + // Append response data for a simple 200 OK + let mut response_data = Vec::new(); + // HEADERS frame with a simple :status 200 + let mut encoder = Encoder::new(4096); + let headers = vec![HeaderField::new(b":status", b"200")]; + let block = encoder.encode(&headers); + let hf = headers_frame(1, block, true); + hf.write_to(&mut response_data).unwrap(); + + // Add response data to the read buffer + conn.stream.read_buf.extend_from_slice(&response_data); + + let stream_id = conn + .send_request("GET", "/", "example.com", &[], None) + .unwrap(); + assert_eq!(stream_id, 1); + assert_eq!(conn.next_stream_id, 3); + + // Read the response + let (_headers, body, status) = conn.read_response(1).unwrap(); + assert_eq!(status, 200); + assert!(body.is_empty()); + } + + #[test] + fn handles_ping() { + let mut server_data = build_server_preface(&[]); + + // Server sends a PING + let ping = ping_frame([1, 2, 3, 4, 5, 6, 7, 8], false); + ping.write_to(&mut server_data).unwrap(); + + // Also send a SETTINGS ACK so we don't hang + // (our WINDOW_UPDATE is already handled) + + let transport = MockTransport::new(server_data); + let mut conn = Http2Connection::new(transport).unwrap(); + + // Process the PING frame + conn.read_and_process_frame().unwrap(); + + // Verify PING ACK was sent (it will be in the write buffer) + // The write buffer contains: preface + settings + window_update + settings_ack + ping_ack + let written = &conn.stream.write_buf; + // Find the PING ACK frame at the end + let last_frame_start = written.len() - 9 - 8; // header + 8 bytes payload + let header = FrameHeader::decode( + &written[last_frame_start..last_frame_start + 9] + .try_into() + .unwrap(), + ); + assert_eq!(header.frame_type, FRAME_PING); + assert!(header.has_flag(FLAG_ACK)); + assert_eq!(&written[last_frame_start + 9..], &[1, 2, 3, 4, 5, 6, 7, 8]); + } + + #[test] + fn handles_goaway_no_error() { + let mut server_data = build_server_preface(&[]); + + // Server sends GOAWAY with NO_ERROR + let ga = goaway_frame(0, ErrorCode::NoError); + ga.write_to(&mut server_data).unwrap(); + + let transport = MockTransport::new(server_data); + let mut conn = Http2Connection::new(transport).unwrap(); + + conn.read_and_process_frame().unwrap(); + assert!(conn.goaway_received); + } + + #[test] + fn handles_goaway_with_error() { + let mut server_data = build_server_preface(&[]); + + let ga = goaway_frame(0, ErrorCode::ProtocolError); + ga.write_to(&mut server_data).unwrap(); + + let transport = MockTransport::new(server_data); + let mut conn = Http2Connection::new(transport).unwrap(); + + let result = conn.read_and_process_frame(); + assert!(matches!( + result, + Err(Http2Error::GoAway(ErrorCode::ProtocolError, _)) + )); + } + + #[test] + fn settings_updates_peer_config() { + let server_data = build_server_preface(&[ + (SETTINGS_MAX_FRAME_SIZE, 32768), + (SETTINGS_HEADER_TABLE_SIZE, 8192), + ]); + + let transport = MockTransport::new(server_data); + let conn = Http2Connection::new(transport).unwrap(); + + assert_eq!(conn.peer_settings.max_frame_size, 32768); + assert_eq!(conn.peer_settings.header_table_size, 8192); + } + + #[test] + fn window_update_connection_level() { + let mut server_data = build_server_preface(&[]); + + // Server sends connection-level WINDOW_UPDATE + let wu = window_update_frame(0, 1000); + wu.write_to(&mut server_data).unwrap(); + + let transport = MockTransport::new(server_data); + let mut conn = Http2Connection::new(transport).unwrap(); + + let initial_window = conn.conn_send_window; + conn.read_and_process_frame().unwrap(); + assert_eq!(conn.conn_send_window, initial_window + 1000); + } + + #[test] + fn full_request_response() { + let mut server_data = build_server_preface(&[]); + + // Build a response: HEADERS with :status 200, then DATA with body + let mut encoder = Encoder::new(4096); + let resp_headers = vec![ + HeaderField::new(b":status", b"200"), + HeaderField::new(b"content-type", b"text/plain"), + ]; + let block = encoder.encode(&resp_headers); + let hf = headers_frame(1, block, false); + hf.write_to(&mut server_data).unwrap(); + + // DATA frame with body + let body = b"Hello, HTTP/2!"; + let df = data_frame(1, body.to_vec(), true); + df.write_to(&mut server_data).unwrap(); + + let transport = MockTransport::new(server_data); + let mut conn = Http2Connection::new(transport).unwrap(); + + let stream_id = conn + .send_request("GET", "/hello", "example.com", &[], None) + .unwrap(); + assert_eq!(stream_id, 1); + + let (headers, resp_body, status) = conn.read_response(1).unwrap(); + assert_eq!(status, 200); + assert_eq!(resp_body, b"Hello, HTTP/2!"); + + // Check headers include content-type + let ct = headers + .iter() + .find(|(n, _)| n == b"content-type") + .map(|(_, v)| v.as_slice()); + assert_eq!(ct, Some(b"text/plain".as_slice())); + } +} diff --git a/crates/net/src/http2/frame.rs b/crates/net/src/http2/frame.rs new file mode 100644 index 0000000..99b6626 --- /dev/null +++ b/crates/net/src/http2/frame.rs @@ -0,0 +1,766 @@ +//! HTTP/2 binary framing layer (RFC 7540 §4). +//! +//! Implements the 9-byte frame header and all standard frame types: +//! DATA, HEADERS, PRIORITY, RST_STREAM, SETTINGS, PUSH_PROMISE, +//! PING, GOAWAY, WINDOW_UPDATE, CONTINUATION. + +use std::fmt; +use std::io::{self, Read, Write}; + +// --------------------------------------------------------------------------- +// Constants +// --------------------------------------------------------------------------- + +/// Size of the HTTP/2 frame header (9 bytes). +pub const FRAME_HEADER_SIZE: usize = 9; + +/// Default maximum frame payload size (RFC 7540 §4.2). +pub const DEFAULT_MAX_FRAME_SIZE: u32 = 16384; + +/// Maximum allowed value for SETTINGS_MAX_FRAME_SIZE (RFC 7540 §4.2). +pub const MAX_FRAME_SIZE_LIMIT: u32 = 16_777_215; + +/// HTTP/2 connection preface magic bytes. +pub const CONNECTION_PREFACE: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n"; + +// --------------------------------------------------------------------------- +// Frame type codes (RFC 7540 §6) +// --------------------------------------------------------------------------- + +pub const FRAME_DATA: u8 = 0x0; +pub const FRAME_HEADERS: u8 = 0x1; +pub const FRAME_PRIORITY: u8 = 0x2; +pub const FRAME_RST_STREAM: u8 = 0x3; +pub const FRAME_SETTINGS: u8 = 0x4; +pub const FRAME_PUSH_PROMISE: u8 = 0x5; +pub const FRAME_PING: u8 = 0x6; +pub const FRAME_GOAWAY: u8 = 0x7; +pub const FRAME_WINDOW_UPDATE: u8 = 0x8; +pub const FRAME_CONTINUATION: u8 = 0x9; + +// --------------------------------------------------------------------------- +// Frame flags +// --------------------------------------------------------------------------- + +pub const FLAG_END_STREAM: u8 = 0x1; +pub const FLAG_END_HEADERS: u8 = 0x4; +pub const FLAG_PADDED: u8 = 0x8; +pub const FLAG_PRIORITY: u8 = 0x20; +pub const FLAG_ACK: u8 = 0x1; + +// --------------------------------------------------------------------------- +// Settings identifiers (RFC 7540 §6.5.2) +// --------------------------------------------------------------------------- + +pub const SETTINGS_HEADER_TABLE_SIZE: u16 = 0x1; +pub const SETTINGS_ENABLE_PUSH: u16 = 0x2; +pub const SETTINGS_MAX_CONCURRENT_STREAMS: u16 = 0x3; +pub const SETTINGS_INITIAL_WINDOW_SIZE: u16 = 0x4; +pub const SETTINGS_MAX_FRAME_SIZE: u16 = 0x5; +pub const SETTINGS_MAX_HEADER_LIST_SIZE: u16 = 0x6; + +// --------------------------------------------------------------------------- +// Error codes (RFC 7540 §7) +// --------------------------------------------------------------------------- + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ErrorCode { + NoError, + ProtocolError, + InternalError, + FlowControlError, + SettingsTimeout, + StreamClosed, + FrameSizeError, + RefusedStream, + Cancel, + CompressionError, + ConnectError, + EnhanceYourCalm, + InadequateSecurity, + Http11Required, +} + +impl ErrorCode { + pub fn from_u32(code: u32) -> Self { + match code { + 0x0 => Self::NoError, + 0x1 => Self::ProtocolError, + 0x2 => Self::InternalError, + 0x3 => Self::FlowControlError, + 0x4 => Self::SettingsTimeout, + 0x5 => Self::StreamClosed, + 0x6 => Self::FrameSizeError, + 0x7 => Self::RefusedStream, + 0x8 => Self::Cancel, + 0x9 => Self::CompressionError, + 0xa => Self::ConnectError, + 0xb => Self::EnhanceYourCalm, + 0xc => Self::InadequateSecurity, + 0xd => Self::Http11Required, + _ => Self::InternalError, + } + } + + pub fn as_u32(self) -> u32 { + match self { + Self::NoError => 0x0, + Self::ProtocolError => 0x1, + Self::InternalError => 0x2, + Self::FlowControlError => 0x3, + Self::SettingsTimeout => 0x4, + Self::StreamClosed => 0x5, + Self::FrameSizeError => 0x6, + Self::RefusedStream => 0x7, + Self::Cancel => 0x8, + Self::CompressionError => 0x9, + Self::ConnectError => 0xa, + Self::EnhanceYourCalm => 0xb, + Self::InadequateSecurity => 0xc, + Self::Http11Required => 0xd, + } + } +} + +impl fmt::Display for ErrorCode { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::NoError => write!(f, "NO_ERROR"), + Self::ProtocolError => write!(f, "PROTOCOL_ERROR"), + Self::InternalError => write!(f, "INTERNAL_ERROR"), + Self::FlowControlError => write!(f, "FLOW_CONTROL_ERROR"), + Self::SettingsTimeout => write!(f, "SETTINGS_TIMEOUT"), + Self::StreamClosed => write!(f, "STREAM_CLOSED"), + Self::FrameSizeError => write!(f, "FRAME_SIZE_ERROR"), + Self::RefusedStream => write!(f, "REFUSED_STREAM"), + Self::Cancel => write!(f, "CANCEL"), + Self::CompressionError => write!(f, "COMPRESSION_ERROR"), + Self::ConnectError => write!(f, "CONNECT_ERROR"), + Self::EnhanceYourCalm => write!(f, "ENHANCE_YOUR_CALM"), + Self::InadequateSecurity => write!(f, "INADEQUATE_SECURITY"), + Self::Http11Required => write!(f, "HTTP_1_1_REQUIRED"), + } + } +} + +// --------------------------------------------------------------------------- +// Error type +// --------------------------------------------------------------------------- + +#[derive(Debug)] +pub enum Http2Error { + /// I/O error during frame read/write. + Io(io::Error), + /// Frame exceeds maximum allowed size. + FrameTooLarge(u32), + /// Protocol violation. + Protocol(String), + /// Stream was reset by remote. + StreamReset(ErrorCode), + /// Connection-level GOAWAY received. + GoAway(ErrorCode, String), + /// Flow control window exceeded. + FlowControl, + /// HPACK header compression/decompression error. + Compression(super::hpack::HpackError), + /// Connection closed. + ConnectionClosed, +} + +impl fmt::Display for Http2Error { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Io(e) => write!(f, "I/O error: {e}"), + Self::FrameTooLarge(sz) => write!(f, "frame too large: {sz} bytes"), + Self::Protocol(msg) => write!(f, "protocol error: {msg}"), + Self::StreamReset(code) => write!(f, "stream reset: {code}"), + Self::GoAway(code, msg) => write!(f, "GOAWAY: {code} {msg}"), + Self::FlowControl => write!(f, "flow control error"), + Self::Compression(e) => write!(f, "HPACK error: {e:?}"), + Self::ConnectionClosed => write!(f, "connection closed"), + } + } +} + +impl From for Http2Error { + fn from(err: io::Error) -> Self { + Self::Io(err) + } +} + +impl From for Http2Error { + fn from(err: super::hpack::HpackError) -> Self { + Self::Compression(err) + } +} + +pub type Result = std::result::Result; + +// --------------------------------------------------------------------------- +// Frame header +// --------------------------------------------------------------------------- + +/// Parsed HTTP/2 frame header (9 bytes on the wire). +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FrameHeader { + /// Payload length (24-bit, max 16,777,215). + pub length: u32, + /// Frame type (DATA=0, HEADERS=1, etc.). + pub frame_type: u8, + /// Frame flags. + pub flags: u8, + /// Stream identifier (31-bit, high bit reserved). + pub stream_id: u32, +} + +impl FrameHeader { + /// Encode the header into 9 bytes. + pub fn encode(&self) -> [u8; 9] { + let mut buf = [0u8; 9]; + buf[0] = (self.length >> 16) as u8; + buf[1] = (self.length >> 8) as u8; + buf[2] = self.length as u8; + buf[3] = self.frame_type; + buf[4] = self.flags; + let sid = self.stream_id & 0x7FFF_FFFF; // mask reserved bit + buf[5] = (sid >> 24) as u8; + buf[6] = (sid >> 16) as u8; + buf[7] = (sid >> 8) as u8; + buf[8] = sid as u8; + buf + } + + /// Decode a frame header from 9 bytes. + pub fn decode(buf: &[u8; 9]) -> Self { + let length = (buf[0] as u32) << 16 | (buf[1] as u32) << 8 | buf[2] as u32; + let frame_type = buf[3]; + let flags = buf[4]; + let stream_id = + (buf[5] as u32) << 24 | (buf[6] as u32) << 16 | (buf[7] as u32) << 8 | buf[8] as u32; + let stream_id = stream_id & 0x7FFF_FFFF; // mask reserved bit + Self { + length, + frame_type, + flags, + stream_id, + } + } + + pub fn has_flag(&self, flag: u8) -> bool { + self.flags & flag != 0 + } +} + +// --------------------------------------------------------------------------- +// Frame (header + payload) +// --------------------------------------------------------------------------- + +/// A complete HTTP/2 frame. +#[derive(Debug, Clone)] +pub struct Frame { + pub header: FrameHeader, + pub payload: Vec, +} + +impl Frame { + pub fn new(frame_type: u8, flags: u8, stream_id: u32, payload: Vec) -> Self { + Self { + header: FrameHeader { + length: payload.len() as u32, + frame_type, + flags, + stream_id, + }, + payload, + } + } + + /// Read a frame from a reader. + pub fn read_from(reader: &mut R, max_frame_size: u32) -> Result { + let mut header_buf = [0u8; 9]; + read_exact(reader, &mut header_buf)?; + let header = FrameHeader::decode(&header_buf); + + if header.length > max_frame_size { + return Err(Http2Error::FrameTooLarge(header.length)); + } + + let mut payload = vec![0u8; header.length as usize]; + if !payload.is_empty() { + read_exact(reader, &mut payload)?; + } + + Ok(Self { header, payload }) + } + + /// Write a frame to a writer. + pub fn write_to(&self, writer: &mut W) -> Result<()> { + let header_bytes = self.header.encode(); + writer.write_all(&header_bytes)?; + if !self.payload.is_empty() { + writer.write_all(&self.payload)?; + } + Ok(()) + } +} + +/// Read exactly `buf.len()` bytes, returning ConnectionClosed on EOF. +fn read_exact(reader: &mut R, buf: &mut [u8]) -> Result<()> { + let mut offset = 0; + while offset < buf.len() { + match reader.read(&mut buf[offset..]) { + Ok(0) => return Err(Http2Error::ConnectionClosed), + Ok(n) => offset += n, + Err(e) if e.kind() == io::ErrorKind::Interrupted => continue, + Err(e) => return Err(Http2Error::Io(e)), + } + } + Ok(()) +} + +// --------------------------------------------------------------------------- +// Frame constructors +// --------------------------------------------------------------------------- + +/// Build a SETTINGS frame. +pub fn settings_frame(settings: &[(u16, u32)], ack: bool) -> Frame { + let flags = if ack { FLAG_ACK } else { 0 }; + let mut payload = Vec::with_capacity(settings.len() * 6); + if !ack { + for &(id, value) in settings { + payload.extend_from_slice(&id.to_be_bytes()); + payload.extend_from_slice(&value.to_be_bytes()); + } + } + Frame::new(FRAME_SETTINGS, flags, 0, payload) +} + +/// Build a WINDOW_UPDATE frame. +pub fn window_update_frame(stream_id: u32, increment: u32) -> Frame { + let payload = (increment & 0x7FFF_FFFF).to_be_bytes().to_vec(); + Frame::new(FRAME_WINDOW_UPDATE, 0, stream_id, payload) +} + +/// Build a HEADERS frame with encoded header block. +pub fn headers_frame(stream_id: u32, header_block: Vec, end_stream: bool) -> Frame { + let mut flags = FLAG_END_HEADERS; + if end_stream { + flags |= FLAG_END_STREAM; + } + Frame::new(FRAME_HEADERS, flags, stream_id, header_block) +} + +/// Build a DATA frame. +pub fn data_frame(stream_id: u32, data: Vec, end_stream: bool) -> Frame { + let flags = if end_stream { FLAG_END_STREAM } else { 0 }; + Frame::new(FRAME_DATA, flags, stream_id, data) +} + +/// Build a RST_STREAM frame. +pub fn rst_stream_frame(stream_id: u32, error_code: ErrorCode) -> Frame { + let payload = error_code.as_u32().to_be_bytes().to_vec(); + Frame::new(FRAME_RST_STREAM, 0, stream_id, payload) +} + +/// Build a GOAWAY frame. +pub fn goaway_frame(last_stream_id: u32, error_code: ErrorCode) -> Frame { + let mut payload = Vec::with_capacity(8); + payload.extend_from_slice(&(last_stream_id & 0x7FFF_FFFF).to_be_bytes()); + payload.extend_from_slice(&error_code.as_u32().to_be_bytes()); + Frame::new(FRAME_GOAWAY, 0, 0, payload) +} + +/// Build a PING frame. +pub fn ping_frame(data: [u8; 8], ack: bool) -> Frame { + let flags = if ack { FLAG_ACK } else { 0 }; + Frame::new(FRAME_PING, flags, 0, data.to_vec()) +} + +/// Build a SETTINGS ACK frame. +pub fn settings_ack_frame() -> Frame { + settings_frame(&[], true) +} + +// --------------------------------------------------------------------------- +// Settings parsing +// --------------------------------------------------------------------------- + +/// A single SETTINGS parameter (id, value). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct Setting { + pub id: u16, + pub value: u32, +} + +/// Parse SETTINGS frame payload into a list of settings. +pub fn parse_settings(payload: &[u8]) -> Result> { + if !payload.len().is_multiple_of(6) { + return Err(Http2Error::Protocol( + "SETTINGS payload length not multiple of 6".to_string(), + )); + } + let mut settings = Vec::with_capacity(payload.len() / 6); + let mut offset = 0; + while offset + 6 <= payload.len() { + let id = u16::from_be_bytes([payload[offset], payload[offset + 1]]); + let value = u32::from_be_bytes([ + payload[offset + 2], + payload[offset + 3], + payload[offset + 4], + payload[offset + 5], + ]); + settings.push(Setting { id, value }); + offset += 6; + } + Ok(settings) +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + // -- FrameHeader encode/decode -- + + #[test] + fn frame_header_roundtrip() { + let header = FrameHeader { + length: 1024, + frame_type: FRAME_DATA, + flags: FLAG_END_STREAM, + stream_id: 1, + }; + let encoded = header.encode(); + let decoded = FrameHeader::decode(&encoded); + assert_eq!(header, decoded); + } + + #[test] + fn frame_header_max_length() { + let header = FrameHeader { + length: MAX_FRAME_SIZE_LIMIT, + frame_type: FRAME_DATA, + flags: 0, + stream_id: 0, + }; + let encoded = header.encode(); + let decoded = FrameHeader::decode(&encoded); + assert_eq!(decoded.length, MAX_FRAME_SIZE_LIMIT); + } + + #[test] + fn frame_header_masks_reserved_bit() { + let mut buf = [0u8; 9]; + // Set the reserved bit (high bit of stream ID) + buf[5] = 0x80; + buf[6] = 0x00; + buf[7] = 0x00; + buf[8] = 0x01; + let header = FrameHeader::decode(&buf); + assert_eq!(header.stream_id, 1); + } + + #[test] + fn frame_header_zero_stream() { + let header = FrameHeader { + length: 0, + frame_type: FRAME_SETTINGS, + flags: 0, + stream_id: 0, + }; + let encoded = header.encode(); + let decoded = FrameHeader::decode(&encoded); + assert_eq!(decoded.stream_id, 0); + } + + #[test] + fn frame_header_has_flag() { + let header = FrameHeader { + length: 0, + frame_type: FRAME_HEADERS, + flags: FLAG_END_STREAM | FLAG_END_HEADERS, + stream_id: 1, + }; + assert!(header.has_flag(FLAG_END_STREAM)); + assert!(header.has_flag(FLAG_END_HEADERS)); + assert!(!header.has_flag(FLAG_PADDED)); + } + + // -- Frame read/write roundtrip -- + + #[test] + fn frame_write_read_roundtrip() { + let frame = Frame::new(FRAME_DATA, FLAG_END_STREAM, 3, vec![1, 2, 3, 4]); + let mut buf = Vec::new(); + frame.write_to(&mut buf).unwrap(); + + let mut cursor = &buf[..]; + let read_frame = Frame::read_from(&mut cursor, DEFAULT_MAX_FRAME_SIZE).unwrap(); + + assert_eq!(read_frame.header, frame.header); + assert_eq!(read_frame.payload, frame.payload); + } + + #[test] + fn frame_empty_payload() { + let frame = Frame::new(FRAME_SETTINGS, FLAG_ACK, 0, vec![]); + let mut buf = Vec::new(); + frame.write_to(&mut buf).unwrap(); + assert_eq!(buf.len(), FRAME_HEADER_SIZE); + + let mut cursor = &buf[..]; + let read_frame = Frame::read_from(&mut cursor, DEFAULT_MAX_FRAME_SIZE).unwrap(); + assert!(read_frame.payload.is_empty()); + } + + #[test] + fn frame_too_large_rejected() { + let header = FrameHeader { + length: DEFAULT_MAX_FRAME_SIZE + 1, + frame_type: FRAME_DATA, + flags: 0, + stream_id: 1, + }; + let mut buf = Vec::new(); + buf.extend_from_slice(&header.encode()); + // Don't need actual payload - error happens before reading it + buf.extend_from_slice(&vec![0u8; (DEFAULT_MAX_FRAME_SIZE + 1) as usize]); + + let mut cursor = &buf[..]; + let result = Frame::read_from(&mut cursor, DEFAULT_MAX_FRAME_SIZE); + assert!(matches!(result, Err(Http2Error::FrameTooLarge(_)))); + } + + // -- Settings frame -- + + #[test] + fn settings_frame_encoding() { + let settings = vec![ + (SETTINGS_MAX_CONCURRENT_STREAMS, 100), + (SETTINGS_INITIAL_WINDOW_SIZE, 65535), + ]; + let frame = settings_frame(&settings, false); + assert_eq!(frame.header.frame_type, FRAME_SETTINGS); + assert_eq!(frame.header.stream_id, 0); + assert_eq!(frame.header.flags, 0); + assert_eq!(frame.payload.len(), 12); // 2 settings * 6 bytes + } + + #[test] + fn settings_ack_empty() { + let frame = settings_ack_frame(); + assert_eq!(frame.header.frame_type, FRAME_SETTINGS); + assert!(frame.header.has_flag(FLAG_ACK)); + assert!(frame.payload.is_empty()); + } + + #[test] + fn parse_settings_roundtrip() { + let settings = vec![ + (SETTINGS_HEADER_TABLE_SIZE, 4096), + (SETTINGS_MAX_CONCURRENT_STREAMS, 128), + (SETTINGS_INITIAL_WINDOW_SIZE, 1048576), + (SETTINGS_MAX_FRAME_SIZE, 32768), + ]; + let frame = settings_frame(&settings, false); + let parsed = parse_settings(&frame.payload).unwrap(); + assert_eq!(parsed.len(), 4); + assert_eq!( + parsed[0], + Setting { + id: SETTINGS_HEADER_TABLE_SIZE, + value: 4096 + } + ); + assert_eq!( + parsed[1], + Setting { + id: SETTINGS_MAX_CONCURRENT_STREAMS, + value: 128 + } + ); + assert_eq!( + parsed[2], + Setting { + id: SETTINGS_INITIAL_WINDOW_SIZE, + value: 1048576 + } + ); + assert_eq!( + parsed[3], + Setting { + id: SETTINGS_MAX_FRAME_SIZE, + value: 32768 + } + ); + } + + #[test] + fn parse_settings_invalid_length() { + let result = parse_settings(&[0, 1, 2, 3, 4]); + assert!(result.is_err()); + } + + // -- WINDOW_UPDATE frame -- + + #[test] + fn window_update_encoding() { + let frame = window_update_frame(1, 32768); + assert_eq!(frame.header.frame_type, FRAME_WINDOW_UPDATE); + assert_eq!(frame.header.stream_id, 1); + assert_eq!(frame.payload.len(), 4); + let increment = u32::from_be_bytes([ + frame.payload[0], + frame.payload[1], + frame.payload[2], + frame.payload[3], + ]); + assert_eq!(increment, 32768); + } + + #[test] + fn window_update_masks_reserved_bit() { + let frame = window_update_frame(0, 0x8000_0001); + let increment = u32::from_be_bytes([ + frame.payload[0], + frame.payload[1], + frame.payload[2], + frame.payload[3], + ]); + // High bit should be masked off + assert_eq!(increment, 1); + } + + // -- HEADERS frame -- + + #[test] + fn headers_frame_with_end_stream() { + let block = vec![0x82, 0x86]; // HPACK-encoded + let frame = headers_frame(1, block.clone(), true); + assert_eq!(frame.header.frame_type, FRAME_HEADERS); + assert!(frame.header.has_flag(FLAG_END_HEADERS)); + assert!(frame.header.has_flag(FLAG_END_STREAM)); + assert_eq!(frame.payload, block); + } + + #[test] + fn headers_frame_without_end_stream() { + let frame = headers_frame(3, vec![0x82], false); + assert!(frame.header.has_flag(FLAG_END_HEADERS)); + assert!(!frame.header.has_flag(FLAG_END_STREAM)); + } + + // -- DATA frame -- + + #[test] + fn data_frame_end_stream() { + let frame = data_frame(1, vec![1, 2, 3], true); + assert_eq!(frame.header.frame_type, FRAME_DATA); + assert!(frame.header.has_flag(FLAG_END_STREAM)); + assert_eq!(frame.payload, vec![1, 2, 3]); + } + + #[test] + fn data_frame_no_end_stream() { + let frame = data_frame(1, vec![1, 2, 3], false); + assert!(!frame.header.has_flag(FLAG_END_STREAM)); + } + + // -- RST_STREAM frame -- + + #[test] + fn rst_stream_encoding() { + let frame = rst_stream_frame(5, ErrorCode::Cancel); + assert_eq!(frame.header.frame_type, FRAME_RST_STREAM); + assert_eq!(frame.header.stream_id, 5); + let code = u32::from_be_bytes([ + frame.payload[0], + frame.payload[1], + frame.payload[2], + frame.payload[3], + ]); + assert_eq!(code, ErrorCode::Cancel.as_u32()); + } + + // -- GOAWAY frame -- + + #[test] + fn goaway_encoding() { + let frame = goaway_frame(7, ErrorCode::NoError); + assert_eq!(frame.header.frame_type, FRAME_GOAWAY); + assert_eq!(frame.header.stream_id, 0); + let last_id = u32::from_be_bytes([ + frame.payload[0], + frame.payload[1], + frame.payload[2], + frame.payload[3], + ]); + let code = u32::from_be_bytes([ + frame.payload[4], + frame.payload[5], + frame.payload[6], + frame.payload[7], + ]); + assert_eq!(last_id, 7); + assert_eq!(code, 0); + } + + // -- PING frame -- + + #[test] + fn ping_frame_encoding() { + let data = [1, 2, 3, 4, 5, 6, 7, 8]; + let frame = ping_frame(data, false); + assert_eq!(frame.header.frame_type, FRAME_PING); + assert!(!frame.header.has_flag(FLAG_ACK)); + assert_eq!(frame.payload, data); + } + + #[test] + fn ping_ack_frame() { + let data = [0u8; 8]; + let frame = ping_frame(data, true); + assert!(frame.header.has_flag(FLAG_ACK)); + } + + // -- ErrorCode -- + + #[test] + fn error_code_roundtrip() { + let codes = [ + ErrorCode::NoError, + ErrorCode::ProtocolError, + ErrorCode::InternalError, + ErrorCode::FlowControlError, + ErrorCode::SettingsTimeout, + ErrorCode::StreamClosed, + ErrorCode::FrameSizeError, + ErrorCode::RefusedStream, + ErrorCode::Cancel, + ErrorCode::CompressionError, + ErrorCode::ConnectError, + ErrorCode::EnhanceYourCalm, + ErrorCode::InadequateSecurity, + ErrorCode::Http11Required, + ]; + for code in codes { + assert_eq!(ErrorCode::from_u32(code.as_u32()), code); + } + } + + #[test] + fn error_code_unknown_maps_to_internal() { + assert_eq!(ErrorCode::from_u32(0xFF), ErrorCode::InternalError); + } + + // -- Connection preface -- + + #[test] + fn connection_preface_correct() { + assert_eq!(CONNECTION_PREFACE, b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n"); + assert_eq!(CONNECTION_PREFACE.len(), 24); + } +} diff --git a/crates/net/src/http2/mod.rs b/crates/net/src/http2/mod.rs index 84c3fbc..fe84aa6 100644 --- a/crates/net/src/http2/mod.rs +++ b/crates/net/src/http2/mod.rs @@ -1,3 +1,9 @@ -//! HTTP/2 protocol implementation. +//! HTTP/2 protocol implementation (RFC 7540/9113). +//! +//! Provides binary framing, stream multiplexing, flow control, +//! HPACK header compression, and an HTTP/2 connection abstraction. +pub mod connection; +pub mod frame; pub mod hpack; +pub mod stream; diff --git a/crates/net/src/http2/stream.rs b/crates/net/src/http2/stream.rs new file mode 100644 index 0000000..663d09f --- /dev/null +++ b/crates/net/src/http2/stream.rs @@ -0,0 +1,347 @@ +//! HTTP/2 stream state machine (RFC 7540 §5.1). +//! +//! Tracks per-stream state transitions: idle → open → half-closed → closed. + +use super::frame::{ErrorCode, Http2Error, Result}; + +// --------------------------------------------------------------------------- +// Default flow control window +// --------------------------------------------------------------------------- + +/// Default initial window size per RFC 7540 §6.9.2. +pub const DEFAULT_INITIAL_WINDOW_SIZE: u32 = 65535; + +// --------------------------------------------------------------------------- +// Stream state (RFC 7540 §5.1) +// --------------------------------------------------------------------------- + +/// State of an HTTP/2 stream. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamState { + /// Stream ID has been reserved but no frames sent/received. + Idle, + /// HEADERS sent/received; both sides can send frames. + Open, + /// Local side has sent END_STREAM; can still receive. + HalfClosedLocal, + /// Remote side has sent END_STREAM; can still send. + HalfClosedRemote, + /// Both sides have sent END_STREAM or RST_STREAM received. + Closed, +} + +// --------------------------------------------------------------------------- +// Stream +// --------------------------------------------------------------------------- + +/// An HTTP/2 stream with state, flow control, and buffered data. +pub struct Stream { + /// Stream identifier (odd = client-initiated). + pub id: u32, + /// Current state. + pub state: StreamState, + /// Send flow control window (how many bytes we can still send). + pub send_window: i64, + /// Receive flow control window (how many bytes the peer can still send). + pub recv_window: i64, + /// Accumulated response header block fragments. + pub header_block: Vec, + /// Decoded response headers (populated after END_HEADERS). + pub response_headers: Vec<(Vec, Vec)>, + /// Accumulated response body data. + pub body: Vec, + /// HTTP status code from :status pseudo-header. + pub status_code: Option, +} + +impl Stream { + /// Create a new client-initiated stream. + pub fn new(id: u32, initial_send_window: u32, initial_recv_window: u32) -> Self { + Self { + id, + state: StreamState::Idle, + send_window: initial_send_window as i64, + recv_window: initial_recv_window as i64, + header_block: Vec::new(), + response_headers: Vec::new(), + body: Vec::new(), + status_code: None, + } + } + + /// Transition to Open state (when HEADERS is sent). + pub fn send_headers(&mut self) -> Result<()> { + match self.state { + StreamState::Idle => { + self.state = StreamState::Open; + Ok(()) + } + _ => Err(Http2Error::Protocol(format!( + "cannot send HEADERS in state {:?}", + self.state + ))), + } + } + + /// Handle sending END_STREAM. + pub fn send_end_stream(&mut self) -> Result<()> { + match self.state { + StreamState::Open => { + self.state = StreamState::HalfClosedLocal; + Ok(()) + } + StreamState::HalfClosedRemote => { + self.state = StreamState::Closed; + Ok(()) + } + _ => Err(Http2Error::Protocol(format!( + "cannot send END_STREAM in state {:?}", + self.state + ))), + } + } + + /// Handle receiving END_STREAM from remote. + pub fn recv_end_stream(&mut self) -> Result<()> { + match self.state { + StreamState::Open => { + self.state = StreamState::HalfClosedRemote; + Ok(()) + } + StreamState::HalfClosedLocal => { + self.state = StreamState::Closed; + Ok(()) + } + _ => Err(Http2Error::Protocol(format!( + "received END_STREAM in state {:?}", + self.state + ))), + } + } + + /// Handle receiving RST_STREAM. + pub fn recv_rst_stream(&mut self, error_code: ErrorCode) -> Result<()> { + self.state = StreamState::Closed; + if error_code != ErrorCode::NoError { + return Err(Http2Error::StreamReset(error_code)); + } + Ok(()) + } + + /// Consume bytes from the send window. Returns error if window is exhausted. + pub fn consume_send_window(&mut self, bytes: u32) -> Result<()> { + self.send_window -= bytes as i64; + if self.send_window < 0 { + return Err(Http2Error::FlowControl); + } + Ok(()) + } + + /// Consume bytes from the receive window. + pub fn consume_recv_window(&mut self, bytes: u32) -> Result<()> { + self.recv_window -= bytes as i64; + if self.recv_window < 0 { + return Err(Http2Error::FlowControl); + } + Ok(()) + } + + /// Increase the send window (on WINDOW_UPDATE from peer). + pub fn increase_send_window(&mut self, increment: u32) -> Result<()> { + self.send_window += increment as i64; + // RFC 7540 §6.9.1: window size must not exceed 2^31 - 1 + if self.send_window > 0x7FFF_FFFF { + return Err(Http2Error::FlowControl); + } + Ok(()) + } + + /// Increase the receive window (when we send WINDOW_UPDATE). + pub fn increase_recv_window(&mut self, increment: u32) { + self.recv_window += increment as i64; + } + + /// Check if the stream can receive data. + pub fn can_recv(&self) -> bool { + matches!(self.state, StreamState::Open | StreamState::HalfClosedLocal) + } + + /// Check if the stream is done (fully closed). + pub fn is_closed(&self) -> bool { + self.state == StreamState::Closed + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn new_stream_is_idle() { + let s = Stream::new(1, 65535, 65535); + assert_eq!(s.state, StreamState::Idle); + assert_eq!(s.id, 1); + } + + #[test] + fn idle_to_open_on_send_headers() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + assert_eq!(s.state, StreamState::Open); + } + + #[test] + fn open_to_half_closed_local_on_send_end_stream() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + s.send_end_stream().unwrap(); + assert_eq!(s.state, StreamState::HalfClosedLocal); + } + + #[test] + fn open_to_half_closed_remote_on_recv_end_stream() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + s.recv_end_stream().unwrap(); + assert_eq!(s.state, StreamState::HalfClosedRemote); + } + + #[test] + fn half_closed_local_to_closed_on_recv_end_stream() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + s.send_end_stream().unwrap(); + s.recv_end_stream().unwrap(); + assert_eq!(s.state, StreamState::Closed); + } + + #[test] + fn half_closed_remote_to_closed_on_send_end_stream() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + s.recv_end_stream().unwrap(); + s.send_end_stream().unwrap(); + assert_eq!(s.state, StreamState::Closed); + } + + #[test] + fn rst_stream_closes() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + s.recv_rst_stream(ErrorCode::NoError).unwrap(); + assert_eq!(s.state, StreamState::Closed); + } + + #[test] + fn rst_stream_with_error() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + let result = s.recv_rst_stream(ErrorCode::Cancel); + assert!(matches!( + result, + Err(Http2Error::StreamReset(ErrorCode::Cancel)) + )); + assert_eq!(s.state, StreamState::Closed); + } + + #[test] + fn send_headers_in_non_idle_fails() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + assert!(s.send_headers().is_err()); + } + + #[test] + fn send_end_stream_in_idle_fails() { + let mut s = Stream::new(1, 65535, 65535); + assert!(s.send_end_stream().is_err()); + } + + #[test] + fn recv_end_stream_in_idle_fails() { + let mut s = Stream::new(1, 65535, 65535); + assert!(s.recv_end_stream().is_err()); + } + + // -- Flow control -- + + #[test] + fn consume_send_window() { + let mut s = Stream::new(1, 65535, 65535); + s.consume_send_window(1000).unwrap(); + assert_eq!(s.send_window, 64535); + } + + #[test] + fn consume_send_window_overflow() { + let mut s = Stream::new(1, 100, 65535); + assert!(s.consume_send_window(101).is_err()); + } + + #[test] + fn consume_recv_window() { + let mut s = Stream::new(1, 65535, 65535); + s.consume_recv_window(5000).unwrap(); + assert_eq!(s.recv_window, 60535); + } + + #[test] + fn increase_send_window() { + let mut s = Stream::new(1, 65535, 65535); + s.consume_send_window(1000).unwrap(); + s.increase_send_window(500).unwrap(); + assert_eq!(s.send_window, 65035); + } + + #[test] + fn increase_send_window_overflow() { + let mut s = Stream::new(1, 0x7FFF_FFFF, 65535); + assert!(s.increase_send_window(1).is_err()); + } + + #[test] + fn increase_recv_window() { + let mut s = Stream::new(1, 65535, 65535); + s.consume_recv_window(1000).unwrap(); + s.increase_recv_window(1000); + assert_eq!(s.recv_window, 65535); + } + + #[test] + fn can_recv_in_open_state() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + assert!(s.can_recv()); + } + + #[test] + fn can_recv_in_half_closed_local() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + s.send_end_stream().unwrap(); + assert!(s.can_recv()); + } + + #[test] + fn cannot_recv_in_half_closed_remote() { + let mut s = Stream::new(1, 65535, 65535); + s.send_headers().unwrap(); + s.recv_end_stream().unwrap(); + assert!(!s.can_recv()); + } + + #[test] + fn is_closed() { + let mut s = Stream::new(1, 65535, 65535); + assert!(!s.is_closed()); + s.send_headers().unwrap(); + assert!(!s.is_closed()); + s.send_end_stream().unwrap(); + s.recv_end_stream().unwrap(); + assert!(s.is_closed()); + } +} diff --git a/crates/net/src/tls/handshake.rs b/crates/net/src/tls/handshake.rs index 0772057..5909caa 100644 --- a/crates/net/src/tls/handshake.rs +++ b/crates/net/src/tls/handshake.rs @@ -41,6 +41,7 @@ const HANDSHAKE_FINISHED: u8 = 20; const EXT_SERVER_NAME: u16 = 0; const EXT_SUPPORTED_GROUPS: u16 = 10; const EXT_SIGNATURE_ALGORITHMS: u16 = 13; +const EXT_ALPN: u16 = 16; const EXT_SUPPORTED_VERSIONS: u16 = 43; const EXT_KEY_SHARE: u16 = 51; @@ -210,7 +211,7 @@ fn random_bytes(buf: &mut [u8]) { /// Build a ClientHello handshake message. /// /// Returns (handshake_message, x25519_private_key). -fn build_client_hello(server_name: &str) -> (Vec, [u8; 32]) { +fn build_client_hello(server_name: &str, alpn_protocols: &[&str]) -> (Vec, [u8; 32]) { // Generate X25519 ephemeral keypair let mut private_key = [0u8; 32]; random_bytes(&mut private_key); @@ -248,7 +249,7 @@ fn build_client_hello(server_name: &str) -> (Vec, [u8; 32]) { push_u8(&mut body, 0); // null // Extensions - let extensions = build_extensions(server_name, &public_key); + let extensions = build_extensions(server_name, &public_key, alpn_protocols); push_u16(&mut body, extensions.len() as u16); push_bytes(&mut body, &extensions); @@ -261,7 +262,11 @@ fn build_client_hello(server_name: &str) -> (Vec, [u8; 32]) { (msg, private_key) } -fn build_extensions(server_name: &str, x25519_public: &[u8; 32]) -> Vec { +fn build_extensions( + server_name: &str, + x25519_public: &[u8; 32], + alpn_protocols: &[&str], +) -> Vec { let mut exts = Vec::with_capacity(256); // SNI extension (server_name) @@ -326,6 +331,20 @@ fn build_extensions(server_name: &str, x25519_public: &[u8; 32]) -> Vec { } } + // ALPN extension (RFC 7301) + if !alpn_protocols.is_empty() { + let mut protocol_list = Vec::new(); + for proto in alpn_protocols { + let bytes = proto.as_bytes(); + protocol_list.push(bytes.len() as u8); + protocol_list.extend_from_slice(bytes); + } + push_u16(&mut exts, EXT_ALPN); + push_u16(&mut exts, (2 + protocol_list.len()) as u16); // extension data length + push_u16(&mut exts, protocol_list.len() as u16); // protocol list length + push_bytes(&mut exts, &protocol_list); + } + exts } @@ -416,12 +435,32 @@ fn parse_server_hello(data: &[u8]) -> Result { // Encrypted handshake message parsing // --------------------------------------------------------------------------- -fn parse_encrypted_extensions(data: &[u8]) -> Result<()> { +fn parse_encrypted_extensions(data: &[u8]) -> Result> { let mut offset = 0; - let _extensions_len = read_u16(data, &mut offset)?; - // We don't require any specific encrypted extensions for now. - // Just validate the format is parseable. - Ok(()) + let extensions_len = read_u16(data, &mut offset)? as usize; + let extensions_end = offset + extensions_len; + let mut alpn_protocol = None; + + while offset < extensions_end { + let ext_type = read_u16(data, &mut offset)?; + let ext_len = read_u16(data, &mut offset)? as usize; + let ext_data = read_bytes(data, &mut offset, ext_len)?; + + if ext_type == EXT_ALPN { + // Parse ALPN response: protocol_list_len(2) + protocol_len(1) + protocol + let mut eoff = 0; + let _list_len = read_u16(ext_data, &mut eoff)?; + let proto_len = read_u8(ext_data, &mut eoff)? as usize; + let proto = read_bytes(ext_data, &mut eoff, proto_len)?; + alpn_protocol = Some( + std::str::from_utf8(proto) + .map_err(|_| HandshakeError::Malformed("invalid ALPN protocol UTF-8"))? + .to_string(), + ); + } + } + + Ok(alpn_protocol) } /// Parse a Certificate handshake message (RFC 8446 §4.4.2). @@ -652,6 +691,23 @@ impl TlsStream { } } +impl io::Read for TlsStream { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + TlsStream::read(self, buf).map_err(|e| io::Error::other(e.to_string())) + } +} + +impl io::Write for TlsStream { + fn write(&mut self, buf: &[u8]) -> io::Result { + TlsStream::write(self, buf).map_err(|e| io::Error::other(e.to_string())) + } + + fn flush(&mut self) -> io::Result<()> { + // TLS writes are flushed per record + Ok(()) + } +} + // --------------------------------------------------------------------------- // Handshake state machine // --------------------------------------------------------------------------- @@ -660,10 +716,22 @@ impl TlsStream { /// /// Returns a `TlsStream` ready for application data. pub fn connect(stream: S, server_name: &str) -> Result> { + let (tls_stream, _alpn) = connect_with_alpn(stream, server_name, &[])?; + Ok(tls_stream) +} + +/// Perform a TLS 1.3 handshake with ALPN negotiation. +/// +/// Returns a `TlsStream` and the negotiated ALPN protocol (if any). +pub fn connect_with_alpn( + stream: S, + server_name: &str, + alpn_protocols: &[&str], +) -> Result<(TlsStream, Option)> { let mut record_layer = RecordLayer::new(stream); // Step 1: Build and send ClientHello - let (client_hello_msg, x25519_private) = build_client_hello(server_name); + let (client_hello_msg, x25519_private) = build_client_hello(server_name, alpn_protocols); let ch_record = TlsRecord::new(ContentType::Handshake, client_hello_msg.clone()); record_layer.write_record(&ch_record)?; @@ -702,7 +770,7 @@ pub fn connect(stream: S, server_name: &str) -> Result(stream: S, server_name: &str) -> Result