From 468f7e74737356c5e37be391d476a00f60948267 Mon Sep 17 00:00:00 2001 From: vshakitskiy Date: Sun, 9 Aug 2026 21:34:20 +0300 Subject: [PATCH] setup http2 implementation --- CHANGELOG.md | 11 + gleam.toml | 3 +- manifest.toml | 6 +- src/ewe.gleam | 335 ++- src/ewe/internal/connection.gleam | 16 +- src/ewe/internal/ewe_ffi.erl | 34 +- src/ewe/internal/file.gleam | 41 +- src/ewe/internal/file_ffi.erl | 25 +- src/ewe/internal/handler.gleam | 92 +- src/ewe/internal/http1_ffi.erl | 43 +- src/ewe/internal/http2.gleam | 2310 +++++++++++++++++ src/ewe/internal/http2/body.gleam | 121 + src/ewe/internal/http2/connection.gleam | 146 ++ src/ewe/internal/http2/frame.gleam | 448 ++++ src/ewe/internal/http2/sse.gleam | 102 + src/ewe/internal/http2/stream.gleam | 162 ++ src/ewe/internal/http2_ffi.erl | 103 + test/ewe/internal/http2/body_test.gleam | 183 ++ test/ewe/internal/http2/connection_test.gleam | 1103 ++++++++ test/ewe/internal/http2/frame_test.gleam | 217 ++ 20 files changed, 5402 insertions(+), 99 deletions(-) create mode 100644 src/ewe/internal/http2.gleam create mode 100644 src/ewe/internal/http2/body.gleam create mode 100644 src/ewe/internal/http2/connection.gleam create mode 100644 src/ewe/internal/http2/frame.gleam create mode 100644 src/ewe/internal/http2/sse.gleam create mode 100644 src/ewe/internal/http2/stream.gleam create mode 100644 src/ewe/internal/http2_ffi.erl create mode 100644 test/ewe/internal/http2/body_test.gleam create mode 100644 test/ewe/internal/http2/connection_test.gleam create mode 100644 test/ewe/internal/http2/frame_test.gleam diff --git a/CHANGELOG.md b/CHANGELOG.md index 044db29..e653466 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -46,6 +46,17 @@ - In replacement of the `CloseCode` variants that each carried their own data there is now `CloseReason`, either `NoCloseReason` or a code and a description which `send_close_frame` takes. +- Add HTTP/2 support!!!! WebSockets need extended CONNECT over HTTP/2 which ewe + does not negotiate yet, so for now any `websocket` answers 501 on an HTTP/2 + connection. +- Add `Http2Options`, `default_http2_options` and `with_http2` to set the limits + and timeouts for every HTTP/2 connection. Concurrent streams, window sizes and + their refill marks, frame and header list sizes, the HPACK table size, the + CONTINUATION and header block caps, the Rapid Reset window and threshold, the + handshake, drain and body read timeouts, and the file read threshold. A value + the protocol does not allow is replaced with the default. +- Add `with_client_verification` which requires clients to present a certificate + signed by a given authority and refuses those that do not. ## v4.0.1 - 04.06.2026 diff --git a/gleam.toml b/gleam.toml index 4f21813..f67acde 100644 --- a/gleam.toml +++ b/gleam.toml @@ -8,9 +8,10 @@ repository = { type = "github", user = "vshakitskiy", repo = "ewe" } links = [{ title = "Gleam", href = "https://gleam.run" }] [dependencies] +alpacki = ">= 3.0.1 and < 4.0.0" gleam_stdlib = ">= 1.0.0 and < 2.0.0" # glisten = ">= 9.0.0 and < 10.0.0" -glisten = { git = "https://github.com/vshakitskiy/glisten.git", ref = "63e9a39" } +glisten = { git = "https://github.com/vshakitskiy/glisten.git", ref = "81a6b00" } gleam_otp = ">= 1.1.0 and < 2.0.0" gleam_http = ">= 4.3.0 and < 5.0.0" logging = ">= 1.3.0 and < 2.0.0" diff --git a/manifest.toml b/manifest.toml index 0a8964c..5b2dfc7 100644 --- a/manifest.toml +++ b/manifest.toml @@ -7,6 +7,7 @@ # You should check this file into your source control repository. packages = [ + { name = "alpacki", version = "3.0.1", build_tools = ["gleam"], requirements = ["gleam_stdlib"], otp_app = "alpacki", source = "hex", outer_checksum = "5BCB617A9606E56018790D4F24AB1E06B6D05BC17ACD6686B99E8D4D18B1F5FD" }, { name = "gleam_crypto", version = "1.6.0", build_tools = ["gleam"], requirements = ["gleam_stdlib"], otp_app = "gleam_crypto", source = "hex", outer_checksum = "2DE9E4EF53CF6FEE049D4F765731F7178F7A11AEFAE00EEE63BF7536B354AD3F" }, { name = "gleam_erlang", version = "1.3.0", build_tools = ["gleam"], requirements = ["gleam_stdlib"], otp_app = "gleam_erlang", source = "hex", outer_checksum = "1124AD3AA21143E5AF0FC5CF3D9529F6DB8CA03E43A55711B60B6B7B3874375C" }, { name = "gleam_http", version = "4.3.0", build_tools = ["gleam"], requirements = ["gleam_stdlib"], otp_app = "gleam_http", source = "hex", outer_checksum = "82EA6A717C842456188C190AFB372665EA56CE13D8559BF3B1DD9E40F619EE0C" }, @@ -14,18 +15,19 @@ packages = [ { name = "gleam_otp", version = "1.2.0", build_tools = ["gleam"], requirements = ["gleam_erlang", "gleam_stdlib"], otp_app = "gleam_otp", source = "hex", outer_checksum = "BA6A294E295E428EC1562DC1C11EA7530DCB981E8359134BEABC8493B7B2258E" }, { name = "gleam_stdlib", version = "1.0.3", build_tools = ["gleam"], requirements = [], otp_app = "gleam_stdlib", source = "hex", outer_checksum = "1F543AFBA5D33DA493E6087F4E4C4F20D899411343512686C98A8ABB2963CF22" }, { name = "gleeunit", version = "1.11.0", build_tools = ["gleam"], requirements = ["gleam_stdlib"], otp_app = "gleeunit", source = "hex", outer_checksum = "EC31ABA74256AEA531EDF8169931D775BBB384FED0A8A1BDC4DD9354E3E21826" }, - { name = "glisten", version = "9.0.1", build_tools = ["gleam"], requirements = ["gleam_erlang", "gleam_otp", "gleam_stdlib", "logging"], source = "git", repo = "https://github.com/vshakitskiy/glisten.git", commit = "63e9a39f7dc35526c2f62128f44f871dab245d15" }, + { name = "glisten", version = "9.0.1", build_tools = ["gleam"], requirements = ["gleam_erlang", "gleam_otp", "gleam_stdlib", "logging"], source = "git", repo = "https://github.com/vshakitskiy/glisten.git", commit = "81a6b0005451b61a14dad2f5892d04cbaad3bc5d" }, { name = "logging", version = "1.5.0", build_tools = ["gleam"], requirements = ["gleam_stdlib"], otp_app = "logging", source = "hex", outer_checksum = "BC5F18CE5DD9686100229FE5409BDC3DD5C46D5A7DF2F804AD2D8F0DD6C5060E" }, { name = "websocks", version = "4.0.1", build_tools = ["gleam"], requirements = ["gleam_crypto", "gleam_erlang", "gleam_stdlib"], otp_app = "websocks", source = "hex", outer_checksum = "89B0C31A032CBE28D4C5FB5CC25A8A690669FEBA501C63F99A4E8F5748750B2A" }, ] [requirements] +alpacki = { version = ">= 3.0.1 and < 4.0.0" } gleam_erlang = { version = ">= 1.3.0 and < 2.0.0" } gleam_http = { version = ">= 4.3.0 and < 5.0.0" } gleam_httpc = { version = ">= 5.0.0 and < 6.0.0" } gleam_otp = { version = ">= 1.1.0 and < 2.0.0" } gleam_stdlib = { version = ">= 1.0.0 and < 2.0.0" } gleeunit = { version = ">= 1.0.0 and < 2.0.0" } -glisten = { git = "https://github.com/vshakitskiy/glisten.git", ref = "63e9a39" } +glisten = { git = "https://github.com/vshakitskiy/glisten.git", ref = "81a6b00" } logging = { version = ">= 1.3.0 and < 2.0.0" } websocks = { version = ">= 4.0.1 and < 5.0.0" } diff --git a/src/ewe.gleam b/src/ewe.gleam index 8ff3db1..46eaeb4 100644 --- a/src/ewe.gleam +++ b/src/ewe.gleam @@ -23,6 +23,9 @@ //// "with_tls", //// "with_http1", //// "default_http1_options", +//// "with_http2", +//// "default_http2_options", +//// "with_client_verification", //// "quiet", //// "on_start" //// ] @@ -143,6 +146,10 @@ import ewe/internal/http1/connection as http1 import ewe/internal/http1/encoder import ewe/internal/http1/sse as http1_sse import ewe/internal/http1/websocket as http1_websocket +import ewe/internal/http2/body as http2_body +import ewe/internal/http2/connection as http2 +import ewe/internal/http2/sse as http2_sse +import ewe/internal/http2/stream as http2_stream import ewe/internal/sse import ewe/internal/websocket import gleam/bytes_tree @@ -231,19 +238,21 @@ fn convert_socket_address(address: glisten.SocketAddress) -> SocketAddress { /// Retrieves the client's socket address from the connection. Returns error if /// the socket information is unavailable. pub fn get_client_info(connection: Connection) -> Result(SocketAddress, Nil) { - case connection { - connection.Http1(connection) -> { - let peername = transport.peername(connection.transport, connection.socket) - use info <- result.map(over: peername) - - case info { - socket.TcpSockName(ip_address:, port:) -> - from_internal_options_ip_address(ip_address) - |> TcpSocketAddress(port:) - socket.UnixSockName(path:) -> UnixSocketAddress(path:) - } - } - connection.Http2 -> todo as "HTTP/2 is not implemented yet!" + // An HTTP/2 handler runs in a process that has no access to the socket, so the + // address is resolved once for the connection and carried on every stream. + let peername = case connection { + connection.Http1(connection) -> + transport.peername(connection.transport, connection.socket) + connection.Http2(connection) -> connection.peer + } + + use info <- result.map(over: peername) + + case info { + socket.TcpSockName(ip_address:, port:) -> + from_internal_options_ip_address(ip_address) + |> TcpSocketAddress(port:) + socket.UnixSockName(path:) -> UnixSocketAddress(path:) } } @@ -375,6 +384,177 @@ fn to_internal_http1_options(options: Http1Options) -> http1.Config { ) } +const max_window_size = 2_147_483_647 + +/// The limits and timeouts for every HTTP/2 connection. Build one by updating +/// `default_http2_options`: +/// +/// ```gleam +/// Http2Options(..ewe.default_http2_options(), max_concurrent_streams: Some(100)) +/// ``` +/// +/// Sizes are in bytes and timeouts in milliseconds. A value the protocol does +/// not allow is replaced with the default. +pub type Http2Options { + Http2Options( + /// How many streams a client may have open at once. `None` leaves it + /// unlimited. + max_concurrent_streams: Option(Int), + /// How much response body a stream may have in flight before the client + /// has to allow more. Must be within 0 and 2147483647. + initial_window_size: Int, + /// The largest frame the server accepts. Must be within 16384 and 16777215. + max_frame_size: Int, + /// The largest header list the server accepts. `None` leaves it unlimited. + max_header_list_size: Option(Int), + /// How much HPACK dynamic table the server keeps for decoding. + header_table_size: Int, + /// How many CONTINUATION frames one header sequence may span. + max_continuation_frames: Int, + /// How many bytes of HEADERS and CONTINUATION one header block may total, + /// counted before it is decoded. + max_header_block_bytes: Int, + /// The window over which client stream resets are counted. + rapid_reset_window: Int, + /// How many resets within that window trip a GOAWAY which is what keeps + /// Rapid Reset (CVE-2023-44487) from costing more than it should. + rapid_reset_threshold: Int, + /// How long a connection may sit in the preface and SETTINGS handshake + /// before it is dropped. + handshake_timeout: Int, + /// How long a draining connection waits for its streams to finish after + /// GOAWAY before closing. + drain_timeout: Int, + /// Once a receive window falls to this it is topped straight back up to + /// `recv_window_high_water_mark` rather than trickling small updates. + recv_window_low_water_mark: Int, + /// What a receive window is topped up to. The wider the gap from the low + /// mark the fewer WINDOW_UPDATE round trips a large body costs. + recv_window_high_water_mark: Int, + /// Files at or below this size are read into memory and framed like any + /// other body. Larger ones are streamed from disk instead. + file_read_threshold: Int, + /// How long a single read of a request body waits for the client. + body_read_timeout: Int, + ) +} + +pub fn default_http2_options() -> Http2Options { + let http2.Config( + max_concurrent_streams:, + initial_window_size:, + max_frame_size:, + max_header_list_size:, + header_table_size:, + max_continuation_frames:, + max_header_block_bytes:, + rapid_reset_window_ms:, + rapid_reset_threshold:, + handshake_timeout_ms:, + drain_timeout_ms:, + recv_window_low_water_mark:, + recv_window_high_water_mark:, + file_read_threshold:, + body_read_timeout:, + ) = http2.default_config() + + Http2Options( + max_concurrent_streams:, + initial_window_size:, + max_frame_size:, + max_header_list_size:, + header_table_size:, + max_continuation_frames:, + max_header_block_bytes:, + rapid_reset_window: rapid_reset_window_ms, + rapid_reset_threshold:, + handshake_timeout: handshake_timeout_ms, + drain_timeout: drain_timeout_ms, + recv_window_low_water_mark:, + recv_window_high_water_mark:, + file_read_threshold:, + body_read_timeout:, + ) +} + +/// Anything the protocol rules out would break connections so it is dropped +/// for the default here instead of reaching a peer. +fn to_internal_http2_options(options: Http2Options) -> http2.Config { + let defaults = http2.default_config() + + let max_concurrent_streams = case options.max_concurrent_streams { + Some(limit) if limit <= 0 -> None + limit -> limit + } + + let initial_window_size = case options.initial_window_size { + size if size < 0 || size > max_window_size -> defaults.initial_window_size + size -> size + } + + let max_frame_size = case options.max_frame_size { + size if size < 16_384 || size > 16_777_215 -> defaults.max_frame_size + size -> size + } + + let drain_timeout_ms = case options.drain_timeout { + timeout if timeout <= 0 -> defaults.drain_timeout_ms + timeout -> timeout + } + + let file_read_threshold = case options.file_read_threshold { + threshold if threshold < 0 -> defaults.file_read_threshold + threshold -> threshold + } + + // The marks only mean anything as a pair, so a bad one replaces both. + let #(recv_window_low_water_mark, recv_window_high_water_mark) = case + options.recv_window_low_water_mark, + options.recv_window_high_water_mark + { + low, high if low > 0 && low < high && high <= max_window_size -> #(low, high) + _low, _high -> #( + defaults.recv_window_low_water_mark, + defaults.recv_window_high_water_mark, + ) + } + + http2.Config( + max_concurrent_streams:, + initial_window_size:, + max_frame_size:, + max_header_list_size: options.max_header_list_size, + header_table_size: options.header_table_size, + max_continuation_frames: options.max_continuation_frames, + max_header_block_bytes: options.max_header_block_bytes, + rapid_reset_window_ms: options.rapid_reset_window, + rapid_reset_threshold: options.rapid_reset_threshold, + handshake_timeout_ms: options.handshake_timeout, + drain_timeout_ms:, + recv_window_low_water_mark:, + recv_window_high_water_mark:, + file_read_threshold:, + body_read_timeout: options.body_read_timeout, + ) +} + +/// Which certificate authority a client's certificate has to be signed by. +pub type ClientVerification { + /// Path to a PEM file holding the CA certificate. + CaCertFile(path: String) + /// In-memory DER-encoded CA certificates. + CaCertData(certs: List(BitArray)) +} + +fn to_internal_client_verification( + verification: ClientVerification, +) -> glisten.CaCert { + case verification { + CaCertFile(path:) -> glisten.CaCertFile(path) + CaCertData(certs:) -> glisten.CaCertData(certs) + } +} + /// Contains all server configurations, can be adjusted by different builder /// functions. pub opaque type Builder { @@ -382,7 +562,9 @@ pub opaque type Builder { handler: fn(request.Request(Connection)) -> response.Response(Body), bind_target: BindTarget, tls: Option(Tls), + client_verification: Option(ClientVerification), http1: Http1Options, + http2: Http2Options, listener_name: process.Name(listener.Message), connection_factory_name: process.Name( factory.Message( @@ -412,7 +594,9 @@ pub fn new( handler:, bind_target: TcpBind(interface: "127.0.0.1", port: 3000, ipv6: False), tls: None, + client_verification: None, http1: default_http1_options(), + http2: default_http2_options(), listener_name:, connection_factory_name:, on_start: fn(scheme, address) { @@ -514,6 +698,20 @@ pub fn with_http1(builder: Builder, options: Http1Options) -> Builder { Builder(..builder, http1: options) } +/// Replaces the limits and timeouts applied to HTTP/2 connections. +pub fn with_http2(builder: Builder, options: Http2Options) -> Builder { + Builder(..builder, http2: options) +} + +/// Requires clients to present a certificate signed by the given authority, +/// refusing those that do not. Needs TLS which `with_tls` configures. +pub fn with_client_verification( + builder: Builder, + ca_cert: ClientVerification, +) -> Builder { + Builder(..builder, client_verification: Some(ca_cert)) +} + fn to_internal_body(body: Body) -> connection.Body { case body { Bytes(tree) -> connection.Bytes(tree) @@ -542,6 +740,7 @@ pub fn start( on_init: handler_.on_init( handler, to_internal_http1_options(builder.http1), + to_internal_http2_options(builder.http2), ), loop: handler_.loop, ) @@ -560,6 +759,22 @@ pub fn start( None -> pool } + // h2c needs nothing announced but over TLS a client only knows HTTP/2 is on + // offer if ALPN says so. + let pool = case builder.tls { + Some(_tls) -> glisten.with_http2(pool) + None -> pool + } + + let pool = case builder.client_verification { + Some(ca_cert) -> + glisten.with_client_verification( + pool, + to_internal_client_verification(ca_cert), + ) + None -> pool + } + use started <- result.map(over: case builder.bind_target { TcpBind(interface:, port:, ipv6:) -> { let pool = glisten.bind(pool, interface) @@ -666,7 +881,22 @@ pub fn read_body( request.Request(..req, headers: list.append(req.headers, trailers), body:) |> Ok } - connection.Http2 -> todo as "HTTP/2 is not implemented yet!" + connection.Http2(connection) -> { + use #(body, trailers) <- result.try( + http2_body.read_body(connection, limit) + |> result.map_error(from_internal_http2_body_error), + ) + + request.Request(..req, headers: list.append(req.headers, trailers), body:) + |> Ok + } + } +} + +fn from_internal_http2_body_error(error: http2_body.BodyError) -> BodyError { + case error { + http2_body.BodyTooLarge -> BodyTooLarge + http2_body.InvalidBody -> InvalidBody } } @@ -701,18 +931,53 @@ pub fn read_body_chunk( Error(error) -> Error(from_internal_http1_body_error(error)) } } - connection.Http2 -> todo as "HTTP/2 is not implemented yet!" + connection.Http2(connection) -> { + case http2_body.read_body_chunk(connection, max_chunk_bytes:, limit:) { + Ok(http2_body.Chunk(data, connection)) -> { + let body = connection.Http2(connection) + Ok(Chunk(data, request.set_body(req, body))) + } + Ok(http2_body.Done(trailers)) -> { + let headers = list.append(req.headers, trailers) + Ok(Done(request.Request(..req, headers:, body: Nil))) + } + Error(error) -> Error(from_internal_http2_body_error(error)) + } + } } } -/// Why a write to the client did not go through. TODO: obviously not the string -/// reason variant but this is for later! +// TODO: remove String reason +/// Why a write to the client did not go through. pub type SendError { - SendError(reason: String) + /// The client is gone so nothing further can be written. + ConnectionClosed + /// The client cancelled this HTTP/2 stream while the rest of the connection + /// carries on. Never returned on HTTP/1. + StreamReset + /// The write failed for a reason the socket reported that does not amount to + /// the client having gone. + SocketError(reason: String) +} + +fn from_interrupted(interrupted: http2.Interrupted) -> SendError { + case interrupted { + http2.StreamReset -> StreamReset + // A write is never given a deadline, so the only way one reports a + // timeout is the connection having stopped answering at all. + http2.ConnectionClosed | http2.TimedOut -> ConnectionClosed + } } fn to_send_error(reason: socket.SocketReason) -> SendError { - SendError(socket.reason_to_string(reason)) + case reason { + socket.Closed + | socket.Econnaborted + | socket.Econnreset + | socket.Enotconn + | socket.Epipe -> ConnectionClosed + reason -> SocketError(socket.reason_to_string(reason)) + } } /// A handle for writing a streamed response's body, obtained from @@ -753,7 +1018,10 @@ pub fn send_chunk( encoder.send_chunk(writer, chunk) |> result.map(connection.Http1Writer) |> result.map_error(to_send_error) - connection.Http2Writer -> todo as "HTTP/2 is not implemented yet!" + connection.Http2Writer(writer) -> + http2_stream.send_chunk(writer, chunk) + |> result.map(connection.Http2Writer) + |> result.map_error(from_interrupted) } } @@ -765,7 +1033,9 @@ pub fn finish_chunk( case writer { connection.Http1Writer(writer) -> encoder.finish_chunk(writer, chunk) |> result.map_error(to_send_error) - connection.Http2Writer -> todo as "HTTP/2 is not implemented yet!" + connection.Http2Writer(writer) -> + http2_stream.finish_chunk(writer, chunk) + |> result.map_error(from_interrupted) } } @@ -775,7 +1045,8 @@ pub fn finish_response(writer: ResponseWriter) -> Result(Nil, SendError) { case writer { connection.Http1Writer(writer) -> encoder.finish_response(writer) |> result.map_error(to_send_error) - connection.Http2Writer -> todo as "HTTP/2 is not implemented yet!" + connection.Http2Writer(writer) -> + http2_stream.finish_response(writer) |> result.map_error(from_interrupted) } } @@ -847,7 +1118,8 @@ pub fn send_event( case conn { connection.Http1Sse(conn) -> http1_sse.send(conn, event) |> result.map_error(to_send_error) - connection.Http2Sse -> todo as "HTTP/2 is not implemented yet!" + connection.Http2Sse(conn) -> + http2_sse.send(conn, event) |> result.map_error(from_interrupted) } } @@ -878,7 +1150,7 @@ pub fn sse( let stream = fn(conn) { case conn { connection.Http1Sse(conn) -> http1_sse.run(conn, on_init, step, on_close) - connection.Http2Sse -> todo as "HTTP/2 is not implemented yet!" + connection.Http2Sse(conn) -> http2_sse.run(conn, on_init, step, on_close) } } @@ -1021,7 +1293,6 @@ pub fn send_text_frame( case conn { connection.Http1Websocket(conn) -> http1_websocket.send_text(conn, text) |> result.map_error(to_send_error) - connection.Http2Websocket -> todo as "HTTP/2 is not implemented yet!" } } @@ -1033,7 +1304,6 @@ pub fn send_binary_frame( case conn { connection.Http1Websocket(conn) -> http1_websocket.send_binary(conn, data) |> result.map_error(to_send_error) - connection.Http2Websocket -> todo as "HTTP/2 is not implemented yet!" } } @@ -1046,7 +1316,6 @@ pub fn send_close_frame( let _sent = case conn { connection.Http1Websocket(conn) -> http1_websocket.send_close(conn, to_internal_close_reason(reason)) - connection.Http2Websocket -> todo as "HTTP/2 is not implemented yet!" } WebsocketStop @@ -1089,7 +1358,6 @@ pub fn websocket( case conn { connection.Http1Websocket(conn) -> http1_websocket.run(conn, on_init, step, on_close) - connection.Http2Websocket -> todo as "HTTP/2 is not implemented yet!" } } @@ -1118,7 +1386,16 @@ pub fn websocket( response.set_body(response.new(400), Empty) } } - connection.Http2 -> todo as "HTTP/2 is not implemented yet!" + // WebSockets ride on extended CONNECT over HTTP/2 which ewe does not + // negotiate yet! + connection.Http2(_connection) -> { + logging.log( + logging.Debug, + "Rejected a WebSocket handshake! HTTP/2 connections do not carry WebSockets", + ) + + response.set_body(response.new(501), Empty) + } } } diff --git a/src/ewe/internal/connection.gleam b/src/ewe/internal/connection.gleam index a1c2c47..aa2f6c7 100644 --- a/src/ewe/internal/connection.gleam +++ b/src/ewe/internal/connection.gleam @@ -1,4 +1,5 @@ import ewe/internal/http1/connection as http1 +import ewe/internal/http2/connection as http2 import gleam/bytes_tree import gleam/erlang/process import gleam/option @@ -8,7 +9,7 @@ import websocks pub type Connection { Http1(http1.Connection) - Http2 + Http2(http2.Connection(Body)) } pub type Body { @@ -52,17 +53,18 @@ pub type FileDescriptor pub type ResponseWriter { Http1Writer(http1.ResponseWriter) - Http2Writer + Http2Writer(http2.ResponseWriter(Body)) } pub type SseConnection { Http1Sse(http1.SseConnection) - Http2Sse + Http2Sse(http2.SseConnection(Body)) } +/// HTTP/2 carries WebSockets over extended CONNECT (RFC 8441), which ewe does +/// not negotiate yet. pub type WebsocketConnection { Http1Websocket(http1.WebsocketConnection) - Http2Websocket } pub type Outcome { @@ -72,6 +74,12 @@ pub type Outcome { pub type Message { Timeout + Http2Handshake + Http2Stream(http2.Reply(Body)) + Http2Exit(process.ExitMessage) + Http2Drain + /// A stream process that outstayed the grace it was given after a reset. + Http2StreamClose(pid: process.Pid) } /// Concatenating onto an empty buffer would copy the incoming bytes for diff --git a/src/ewe/internal/ewe_ffi.erl b/src/ewe/internal/ewe_ffi.erl index e696159..333646a 100644 --- a/src/ewe/internal/ewe_ffi.erl +++ b/src/ewe/internal/ewe_ffi.erl @@ -5,9 +5,12 @@ now_datetime/0, set_http_date/1, get_http_date/0, - rescue_handler/1 + rescue_handler/1, + is_valid_utf8/1 ]). +-define(HIGH_BITS, 16#80808080808080). + identity(X) -> X. @@ -49,3 +52,32 @@ ensure_http_date_table() -> _ -> ?MODULE end. + +is_valid_utf8(Bin) when is_binary(Bin) -> + case skip_ascii(Bin) of + <<>> -> true; + Rest -> is_binary(unicode:characters_to_binary(Rest, utf8)) + end; +is_valid_utf8(_Bits) -> + false. + +%% Tests seven bytes per word rather than eight since 56 bits is the widest that +%% still fits an immediate integer on a 64-bit VM so no word allocates. +skip_ascii(<>) when + A band ?HIGH_BITS =:= 0, + B band ?HIGH_BITS =:= 0, + C band ?HIGH_BITS =:= 0, + D band ?HIGH_BITS =:= 0 +-> + skip_ascii(Rest); +skip_ascii(<>) when Word band ?HIGH_BITS =:= 0 -> + skip_ascii(Rest); +%% Tails shorter than a word, each masked to its own width. +skip_ascii(<>) when Word band 16#808080808080 =:= 0 -> <<>>; +skip_ascii(<>) when Word band 16#8080808080 =:= 0 -> <<>>; +skip_ascii(<>) when Word band 16#80808080 =:= 0 -> <<>>; +skip_ascii(<>) when Word band 16#808080 =:= 0 -> <<>>; +skip_ascii(<>) when Word band 16#8080 =:= 0 -> <<>>; +skip_ascii(<>) when Word band 16#80 =:= 0 -> <<>>; +skip_ascii(Rest) -> + Rest. diff --git a/src/ewe/internal/file.gleam b/src/ewe/internal/file.gleam index ee5599c..0dde35b 100644 --- a/src/ewe/internal/file.gleam +++ b/src/ewe/internal/file.gleam @@ -34,7 +34,7 @@ pub fn resolve( } } } - connection.Http2 -> { + connection.Http2(_connection) -> { use size <- result.try(stat(path)) use #(offset, length) <- result.map(range(size, offset, limit)) @@ -92,11 +92,14 @@ pub fn send( case file { connection.OpenFile(handle:, offset:, length:) -> send_handle(transport, socket, handle, offset, length) - connection.PendingFile(..) -> todo as "HTTP/2 is not implemented yet!" + // Only HTTP/2 leaves a file unopened and its connection process opens one + // itself rather than writing it through here. + connection.PendingFile(..) -> + panic as "an unopened file cannot be written to an HTTP/1 socket" } } -/// Owns the descriptor from here on, so it is closed however the write ends. +/// Owns the descriptor from here on so it is closed however the write ends. fn send_handle( transport: transport.Transport, socket: socket.Socket, @@ -104,7 +107,23 @@ fn send_handle( offset: Int, length: Int, ) -> Result(Nil, socket.SocketReason) { - let sent = case length { + let sent = send_chunk(transport, socket, handle, offset, length) + + close(handle) + sent +} + +/// Writes one range of an open file leaving the descriptor open. HTTP/2 sizes +/// each range to the stream's send window and comes back for the next one so +/// the descriptor has to outlive the individual write. +pub fn send_chunk( + transport: transport.Transport, + socket: socket.Socket, + handle: connection.FileDescriptor, + offset: Int, + length: Int, +) -> Result(Nil, socket.SocketReason) { + case length { 0 -> Ok(Nil) _length -> case transport { @@ -112,9 +131,6 @@ fn send_handle( transport.Ssl -> send_chunks(transport, socket, handle, offset, length) } } - - close(handle) - sent } const chunk_size = 65_536 @@ -148,6 +164,13 @@ fn send_chunks( } } +@external(erlang, "file_ffi", "read_range") +pub fn read_range( + path: String, + offset: Int, + length: Int, +) -> Result(BitArray, FileError) + @external(erlang, "file_ffi", "stat") fn stat(path: String) -> Result(Int, FileError) @@ -160,7 +183,7 @@ fn do_sendfile( ) -> Result(Nil, socket.SocketReason) @external(erlang, "file_ffi", "open") -fn open(path: String) -> Result(connection.FileDescriptor, FileError) +pub fn open(path: String) -> Result(connection.FileDescriptor, FileError) @external(erlang, "file_ffi", "size") fn size(handle: connection.FileDescriptor) -> Result(Int, FileError) @@ -173,4 +196,4 @@ fn pread( ) -> Result(BitArray, socket.SocketReason) @external(erlang, "file_ffi", "close") -fn close(handle: connection.FileDescriptor) -> Nil +pub fn close(handle: connection.FileDescriptor) -> Nil diff --git a/src/ewe/internal/file_ffi.erl b/src/ewe/internal/file_ffi.erl index 96894b4..d0a5e1e 100644 --- a/src/ewe/internal/file_ffi.erl +++ b/src/ewe/internal/file_ffi.erl @@ -2,7 +2,7 @@ -include_lib("kernel/include/file.hrl"). --export([stat/1, sendfile/4, open/1, size/1, pread/3, close/1]). +-export([stat/1, sendfile/4, open/1, size/1, pread/3, read_range/3, close/1]). stat(Path) -> case file:read_file_info(Path, [raw, {time, posix}]) of @@ -15,7 +15,8 @@ stat(Path) -> sendfile(Fd, Socket, Offset, Bytes) -> case file:sendfile(Fd, Socket, Offset, Bytes, []) of - {ok, _Sent} -> {ok, nil}; + {ok, Bytes} -> {ok, nil}; + {ok, _Short} -> {error, closed}; {error, Reason} -> {error, Reason} end. @@ -41,6 +42,26 @@ pread(Fd, Offset, Length) -> {error, Reason} -> {error, Reason} end. +read_range(_Path, _Offset, 0) -> + {ok, <<>>}; +%% A short read means the file shrank between being sized and being read which +%% leaves no correct body to send so it is reported rather than padded over. +read_range(Path, Offset, Length) -> + case open(Path) of + {ok, Fd} -> + Result = + case file:pread(Fd, Offset, Length) of + {ok, Data} when byte_size(Data) =:= Length -> {ok, Data}; + {ok, _Short} -> {error, unknown_error}; + eof -> {error, unknown_error}; + {error, _Reason} -> {error, unknown_error} + end, + file:close(Fd), + Result; + {error, Reason} -> + {error, Reason} + end. + close(Fd) -> file:close(Fd), nil. diff --git a/src/ewe/internal/handler.gleam b/src/ewe/internal/handler.gleam index acd7f47..0a68e63 100644 --- a/src/ewe/internal/handler.gleam +++ b/src/ewe/internal/handler.gleam @@ -1,25 +1,30 @@ import ewe/internal/connection import ewe/internal/http1 import ewe/internal/http1/connection as http1_connection +import ewe/internal/http2 +import ewe/internal/http2/connection as http2_connection import gleam/erlang/process import gleam/http/request import gleam/http/response import gleam/option import glisten +import glisten/socket/options +import glisten/transport import logging /// The connection's protocol, still undecided until the HTTP/2 preface has been /// ruled in or out. pub type State { - Initialised(http1.State) + Initialised(http1.State, http2_connection.Config) Http1(http1.State) - Http2 + Http2(http2.State) } pub fn on_init( handler: fn(request.Request(connection.Connection)) -> response.Response(connection.Body), - config: http1_connection.Config, + http1_config: http1_connection.Config, + http2_config: http2_connection.Config, ) { fn(connection: glisten.Connection(connection.Message)) -> #( State, @@ -29,11 +34,14 @@ pub fn on_init( http1.State( handler:, buffer: <<>>, - idle_timer: connection.start_idle_timer(connection, config.idle_timeout), - config:, + idle_timer: connection.start_idle_timer( + connection, + http1_config.idle_timeout, + ), + config: http1_config, ) - #(Initialised(state), option.None) + #(Initialised(state, http2_config), option.None) } } @@ -43,7 +51,7 @@ pub fn loop( connection: glisten.Connection(connection.Message), ) -> glisten.Next(State, glisten.Message(connection.Message)) { case state, message { - Initialised(state), glisten.Packet(data) -> { + Initialised(state, http2_config), glisten.Packet(data) -> { connection.cancel_idle_timer(state.idle_timer) let buffer = connection.append_buffer(state.buffer, data) @@ -57,28 +65,74 @@ pub fn loop( state.config.idle_timeout, ), ) - |> Initialised + |> Initialised(http2_config) |> glisten.continue - Http2Preface(_remaining) -> glisten.continue(Http2) + Http2Preface(remaining:) -> + start_http2(connection, state.handler, http2_config, remaining) NotHttp2(buffer:) -> http1.State(..state, buffer:, idle_timer: option.None) |> http1.handle_message(connection) - |> to_glisten_next + |> from_http1 } } Http1(state), glisten.Packet(data) -> http1.State(..state, buffer: connection.append_buffer(state.buffer, data)) |> http1.handle_message(connection) - |> to_glisten_next - Http2(..), glisten.Packet(_data) -> todo as "HTTP/2 is not implemented yet!" - _state, glisten.User(connection.Timeout) -> { + |> from_http1 + Http2(state), message -> + http2.handle_message(state, message, connection) + |> from_http2 + Initialised(..), glisten.User(connection.Timeout) + | Http1(..), glisten.User(connection.Timeout) + -> { logging.log(logging.Debug, "Connection idled for too long, closing.") glisten.stop() } + Initialised(..), glisten.User(_message) + | Http1(..), glisten.User(_message) + -> glisten.continue(state) } } -fn to_glisten_next( +/// Takes the connection over for HTTP/2. The client's already waiting on our +/// SETTINGS by now, so that goes out first. +fn start_http2( + connection: glisten.Connection(connection.Message), + handler: fn(request.Request(connection.Connection)) -> + response.Response(connection.Body), + config: http2_connection.Config, + remaining: BitArray, +) -> glisten.Next(State, glisten.Message(connection.Message)) { + process.trap_exits(True) + + let self = process.new_subject() + let replies = process.new_subject() + + let peer = transport.peername(connection.transport, connection.socket) + + let state = http2.init(handler, config, self, replies, peer) + + let selector = + process.new_selector() + |> process.select_map(self, glisten.User) + |> process.select_map(replies, fn(reply) { + glisten.User(connection.Http2Stream(reply)) + }) + |> process.select_trapped_exits(fn(exit) { + glisten.User(connection.Http2Exit(exit)) + }) + + case glisten.send(connection, state.settings_frame) { + Error(_reason) -> glisten.stop() + Ok(Nil) -> + http2.handle_message(state, glisten.Packet(remaining), connection) + |> from_http2 + |> glisten.with_selector(selector) + |> glisten.set_active_state(options.Count(http2.socket_active_batch_size)) + } +} + +fn from_http1( next: http1.Next, ) -> glisten.Next(State, glisten.Message(connection.Message)) { case next { @@ -88,6 +142,16 @@ fn to_glisten_next( } } +fn from_http2( + next: http2.Next, +) -> glisten.Next(State, glisten.Message(connection.Message)) { + case next { + http2.Continue(state) -> glisten.continue(Http2(state)) + http2.Close -> glisten.stop() + http2.CloseAbnormal(reason:) -> glisten.stop_abnormal(reason) + } +} + pub type Sniff { NeedMoreData Http2Preface(remaining: BitArray) diff --git a/src/ewe/internal/http1_ffi.erl b/src/ewe/internal/http1_ffi.erl index 2c78516..780e9ec 100644 --- a/src/ewe/internal/http1_ffi.erl +++ b/src/ewe/internal/http1_ffi.erl @@ -15,7 +15,6 @@ socket_error_reason/1 ]). --define(HIGH_BITS, 16#80808080808080). %% Compiles and caches match patterns once at module load. init() -> @@ -67,40 +66,10 @@ lower(Byte) -> Byte. socket_error_reason({_Tag, _Socket, Reason}) -> Reason. -%% Validates UTF-8 and returns the bytes unchanged. Everything skip_ascii walks -%% past is ASCII which is valid UTF-8 and never part of a multi-byte sequence -%% so whatever it stops on still starts on a character boundary and can be -%% validated on its own. -bit_array_to_string(Bin) when is_binary(Bin) -> - case skip_ascii(Bin) of - <<>> -> - {ok, Bin}; - Rest -> - case unicode:characters_to_binary(Rest, utf8) of - Out when is_binary(Out) -> {ok, Bin}; - _Invalid -> {error, nil} - end - end; -bit_array_to_string(_Bits) -> - {error, nil}. +%% Validates UTF-8 and returns the bytes unchanged. +bit_array_to_string(Bin) -> + case ewe_ffi:is_valid_utf8(Bin) of + true -> {ok, Bin}; + false -> {error, nil} + end. -%% Tests seven bytes per word rather than eight, since 56 bits is the widest that -%% still fits an immediate integer on a 64-bit VM so no word allocates. -skip_ascii(<>) when - A band ?HIGH_BITS =:= 0, - B band ?HIGH_BITS =:= 0, - C band ?HIGH_BITS =:= 0, - D band ?HIGH_BITS =:= 0 --> - skip_ascii(Rest); -skip_ascii(<>) when Word band ?HIGH_BITS =:= 0 -> - skip_ascii(Rest); -%% Tails shorter than a word, each masked to its own width. -skip_ascii(<>) when Word band 16#808080808080 =:= 0 -> <<>>; -skip_ascii(<>) when Word band 16#8080808080 =:= 0 -> <<>>; -skip_ascii(<>) when Word band 16#80808080 =:= 0 -> <<>>; -skip_ascii(<>) when Word band 16#808080 =:= 0 -> <<>>; -skip_ascii(<>) when Word band 16#8080 =:= 0 -> <<>>; -skip_ascii(<>) when Word band 16#80 =:= 0 -> <<>>; -skip_ascii(Rest) -> - Rest. diff --git a/src/ewe/internal/http2.gleam b/src/ewe/internal/http2.gleam new file mode 100644 index 0000000..83bb095 --- /dev/null +++ b/src/ewe/internal/http2.gleam @@ -0,0 +1,2310 @@ +import alpacki +import ewe/internal/clock +import ewe/internal/connection +import ewe/internal/file +import ewe/internal/http2/connection as http2 +import ewe/internal/http2/frame +import ewe/internal/http2/stream +import gleam/bit_array +import gleam/bytes_tree +import gleam/dict.{type Dict} +import gleam/erlang/process +import gleam/http +import gleam/http/request.{type Request, Request} +import gleam/http/response.{type Response} +import gleam/int +import gleam/list +import gleam/option.{type Option, None, Some} +import gleam/result +import glisten +import glisten/socket + +@internal +pub type PeerSettings { + PeerSettings( + header_table_size: Int, + initial_window_size: Int, + max_frame_size: Int, + max_header_list_size: Option(Int), + ) +} + +const default_peer_settings = PeerSettings( + header_table_size: 4096, + initial_window_size: 65_535, + max_frame_size: 16_384, + max_header_list_size: None, +) + +fn apply_settings( + settings: PeerSettings, + params: List(frame.Setting), +) -> PeerSettings { + list.fold(params, settings, apply_setting) +} + +fn apply_setting( + settings: PeerSettings, + setting: frame.Setting, +) -> PeerSettings { + case setting { + frame.HeaderTableSize(value) -> + PeerSettings(..settings, header_table_size: value) + frame.InitialWindowSize(value) -> + PeerSettings(..settings, initial_window_size: value) + frame.MaxFrameSize(value) -> PeerSettings(..settings, max_frame_size: value) + frame.MaxHeaderListSize(value) -> + PeerSettings(..settings, max_header_list_size: Some(value)) + frame.EnablePush(_enabled) + | frame.MaxConcurrentStreams(_limit) + | frame.UnknownSetting(_id, _value) -> settings + } +} + +/// What the connection does with the socket once a message has been handled. +pub type Next { + Continue(State) + Close + CloseAbnormal(reason: String) +} + +@internal +pub type HandshakePhase { + AwaitingSettings + Connected +} + +@internal +pub type HeaderAssembly { + HeaderAssembly( + stream_id: Int, + end_stream: Bool, + fragment_count: Int, + block: BitArray, + trailers: Bool, + ) +} + +@internal +pub type StreamStatus { + Computing(pid: process.Pid) + Flushing +} + +@internal +pub type Pending { + PendingBytes(BitArray) + PendingFile( + descriptor: connection.FileDescriptor, + offset: Int, + remaining: Int, + ) +} + +@internal +pub type Stream { + Stream( + status: StreamStatus, + send_window: Int, + pending: Pending, + pending_end_stream: Bool, + write_ack: Option(process.Subject(http2.WriteAck)), + recv_window: Int, + recv_buffer: bytes_tree.BytesTree, + request_half_closed: Bool, + parked_reader: Option(process.Subject(http2.BodyEvent)), + content_length: Option(Int), + body_bytes_received: Int, + trailers: List(#(String, String)), + ) +} + +// Opaque handle to a `binary:compile_pattern/1` result. Compiled once per +// node and fetched once per connection. +type Pattern + +pub opaque type HeaderPatterns { + HeaderPatterns( + name: Pattern, + forbidden: Pattern, + query: Pattern, + colon: Pattern, + ) +} + +@external(erlang, "http2_ffi", "name_pattern") +fn name_pattern() -> Pattern + +@external(erlang, "http2_ffi", "forbidden_header_pattern") +fn forbidden_header_pattern() -> Pattern + +@external(erlang, "http2_ffi", "query_pattern") +fn query_pattern() -> Pattern + +@external(erlang, "http2_ffi", "colon_pattern") +fn colon_pattern() -> Pattern + +@internal +pub fn header_patterns() -> HeaderPatterns { + HeaderPatterns( + name: name_pattern(), + forbidden: forbidden_header_pattern(), + query: query_pattern(), + colon: colon_pattern(), + ) +} + +@internal +pub type State { + State( + buffer: BitArray, + handshake: HandshakePhase, + peer_settings: PeerSettings, + timer: Option(process.Timer), + hpack_decoder: alpacki.DynamicTable, + hpack_encoder: alpacki.DynamicTable, + header_assembly: Option(HeaderAssembly), + reply_subject: process.Subject(http2.Reply(connection.Body)), + handler: fn(Request(connection.Connection)) -> Response(connection.Body), + streams: Dict(Int, Stream), + stream_pids: Dict(process.Pid, Int), + conn_send_window: Int, + conn_recv_window: Int, + reset_window_start: Int, + reset_count: Int, + highest_client_stream_id_seen: Int, + draining: Bool, + drain_subject: process.Subject(connection.Message), + drain_timer: Option(process.Timer), + config: http2.Config, + settings_frame: bytes_tree.BytesTree, + patterns: HeaderPatterns, + peer: Result(socket.SockName, Nil), + ) +} + +const default_send_window = 65_535 + +const max_window_size = 2_147_483_647 + +/// How many socket messages arrive before the connection has to ask for more. +/// Caps how fast a peer can grow the mailbox. +pub const socket_active_batch_size = 32 + +fn build_settings_frame(config: http2.Config) -> bytes_tree.BytesTree { + let params = case config.header_table_size { + 4096 -> [] + _size -> [frame.HeaderTableSize(config.header_table_size)] + } + + let params = case config.initial_window_size { + 65_535 -> params + _size -> [frame.InitialWindowSize(config.initial_window_size), ..params] + } + + let params = case config.max_frame_size { + 16_384 -> params + _size -> [frame.MaxFrameSize(config.max_frame_size), ..params] + } + + let params = case config.max_concurrent_streams { + Some(value) -> [frame.MaxConcurrentStreams(value), ..params] + None -> params + } + + let params = case config.max_header_list_size { + Some(value) -> [frame.MaxHeaderListSize(value), ..params] + None -> params + } + + frame.Settings(0, False, params) + |> frame.encode + |> bytes_tree.from_bit_array +} + +/// Wired to glisten's close callback so a connection dropped underneath its +/// streams still releases what they were holding. +pub fn kill_live_workers(state: State) -> Nil { + use _stream_id, entry <- dict.each(state.streams) + + case entry.status { + Computing(pid) -> process.send_abnormal_exit(pid, "connection_closed") + Flushing -> Nil + } + + case entry.pending { + PendingFile(descriptor, _offset, _remaining) -> file.close(descriptor) + PendingBytes(_bytes) -> Nil + } +} + +/// Builds the state for a connection whose preface is already read. The caller +/// makes the subjects because the caller is what selects on them. +pub fn init( + handler: fn(Request(connection.Connection)) -> Response(connection.Body), + config: http2.Config, + self: process.Subject(connection.Message), + reply_subject: process.Subject(http2.Reply(connection.Body)), + peer: Result(socket.SockName, Nil), +) -> State { + let table = alpacki.new_dynamic(config.header_table_size) + + let timer = + process.send_after( + self, + config.handshake_timeout_ms, + connection.Http2Handshake, + ) + + State( + buffer: <<>>, + handshake: AwaitingSettings, + peer_settings: default_peer_settings, + timer: Some(timer), + hpack_decoder: table, + hpack_encoder: table, + header_assembly: None, + reply_subject:, + handler:, + streams: dict.new(), + stream_pids: dict.new(), + conn_send_window: default_send_window, + conn_recv_window: default_send_window, + reset_window_start: 0, + reset_count: 0, + highest_client_stream_id_seen: 0, + draining: False, + drain_subject: self, + drain_timer: None, + config:, + settings_frame: build_settings_frame(config), + patterns: header_patterns(), + peer:, + ) +} + +pub fn handle_message( + state: State, + message: glisten.Message(connection.Message), + connection: glisten.Connection(connection.Message), +) -> Next { + case message { + glisten.Packet(bytes) -> handle_packet(state, bytes, connection) + glisten.User(connection.Http2Handshake) -> + handle_handshake_timeout(state, connection) + glisten.User(connection.Http2Stream(reply)) -> + handle_stream_reply(state, reply, connection) + glisten.User(connection.Http2Exit(exit)) -> + handle_stream_exit(state, exit, connection) + glisten.User(connection.Http2Drain) -> stop_connection(state) + glisten.User(connection.Http2StreamClose(pid)) -> + finish_or_continue(handle_stream_close_timeout(state, pid)) + // HTTP/1's idle timer is cancelled before a connection becomes HTTP/2 + // which has its own handshake and drain deadlines instead. + glisten.User(connection.Timeout) -> Continue(state) + } +} + +fn stop_connection(state: State) -> Next { + kill_live_workers(state) + + Close +} + +fn handle_handshake_timeout( + state: State, + connection: glisten.Connection(connection.Message), +) -> Next { + case state.handshake { + AwaitingSettings -> + terminate(state, connection, Some(frame.SettingsTimeout)) + Connected -> Continue(state) + } +} + +fn handle_packet( + state: State, + bytes: BitArray, + connection: glisten.Connection(connection.Message), +) -> Next { + State(..state, buffer: <>) + |> process_frames(connection) +} + +fn process_frames( + state: State, + connection: glisten.Connection(connection.Message), +) -> Next { + case frame.decode(state.buffer, state.config.max_frame_size) { + Ok(#(frame, remaining)) -> + case handle_frame(State(..state, buffer: remaining), frame, connection) { + Proceed(state) -> process_frames(state, connection) + ProceedWithOutbound(state, out) -> + case glisten.send(connection, out) { + Ok(Nil) -> process_frames(state, connection) + Error(_reason) -> terminate(state, connection, None) + } + RejectStream(state, stream_id, code) -> + reject_stream(connection, state, stream_id, code) + Terminate(code) -> terminate(state, connection, code) + } + Error(frame.Incomplete) -> finish_or_continue(state) + Error(frame.Violation(code)) -> terminate(state, connection, Some(code)) + } +} + +fn reject_stream( + connection: glisten.Connection(connection.Message), + state: State, + stream_id: Int, + code: frame.ErrorCode, +) -> Next { + let state = case dict.get(state.streams, stream_id) { + Error(Nil) -> state + Ok(entry) -> reset_and_remove_stream(state, stream_id, entry) + } + + case send_frame(connection, frame.RstStream(stream_id, code)) { + Ok(Nil) -> process_frames(state, connection) + Error(_reason) -> terminate(state, connection, None) + } +} + +const stream_close_grace_ms = 5000 + +fn reset_and_remove_stream( + state: State, + stream_id: Int, + entry: Stream, +) -> State { + case entry.status { + Computing(pid) -> { + process.send_abnormal_exit(pid, "stream_reset") + // Backstop for a long-lived handler that doesn't react to the exit + // signal promptly + process.send_after( + state.drain_subject, + stream_close_grace_ms, + connection.Http2StreamClose(pid), + ) + + Nil + } + Flushing -> Nil + } + + case entry.pending { + PendingFile(descriptor, _offset, _remaining) -> file.close(descriptor) + PendingBytes(_bytes) -> Nil + } + + State(..state, streams: dict.delete(state.streams, stream_id)) +} + +fn handle_stream_close_timeout(state: State, pid: process.Pid) -> State { + case dict.get(state.stream_pids, pid) { + Error(Nil) -> state + Ok(_stream_id) -> { + process.kill(pid) + clear_stream_pid(state, pid) + } + } +} + +@internal +pub type FrameResult { + Proceed(state: State) + ProceedWithOutbound(state: State, out: bytes_tree.BytesTree) + RejectStream(state: State, stream_id: Int, code: frame.ErrorCode) + Terminate(code: Option(frame.ErrorCode)) +} + +fn handle_frame( + state: State, + frame: frame.Frame, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + case state.header_assembly { + Some(assembly) -> handle_continuation(state, assembly, frame) + None -> handle_new_frame(state, frame, connection) + } +} + +fn handle_new_frame( + state: State, + frame: frame.Frame, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + case frame { + frame.Settings(0, True, _params) -> Proceed(state) + frame.Settings(0, False, params) -> + handle_client_settings(state, params, connection) + frame.Settings(..) -> Terminate(Some(frame.ProtocolError)) + frame.Headers(stream_id, end_stream, end_headers, payload) -> + handle_headers(state, stream_id, end_stream, end_headers, payload) + frame.Data(stream_id, end_stream, payload, flow_control_size) -> + case state.handshake { + AwaitingSettings -> Terminate(Some(frame.ProtocolError)) + Connected -> + handle_data(state, stream_id, end_stream, payload, flow_control_size) + } + frame.RstStream(stream_id, _error_code) -> + case state.handshake { + AwaitingSettings -> Terminate(Some(frame.ProtocolError)) + Connected -> handle_client_reset(state, stream_id) + } + frame.WindowUpdate(stream_id, increment) -> + case state.handshake { + AwaitingSettings -> Terminate(Some(frame.ProtocolError)) + Connected -> + handle_window_update(state, stream_id, increment, connection) + } + frame.Ping(stream_id, ack, opaque_data) -> + case state.handshake { + AwaitingSettings -> Terminate(Some(frame.ProtocolError)) + Connected -> handle_ping(state, stream_id, ack, opaque_data, connection) + } + frame.Continuation(..) | frame.PushPromise(..) -> + Terminate(Some(frame.ProtocolError)) + frame.Priority(..) | frame.Goaway(..) | frame.Unknown(..) -> + case state.handshake { + AwaitingSettings -> Terminate(Some(frame.ProtocolError)) + Connected -> Proceed(state) + } + } +} + +fn handle_ping( + state: State, + stream_id: Int, + ack: Bool, + opaque_data: BitArray, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + case stream_id != 0, ack { + True, _ack -> Terminate(Some(frame.ProtocolError)) + False, True -> Proceed(state) + False, False -> + case send_frame(connection, frame.Ping(0, True, opaque_data)) { + Ok(Nil) -> Proceed(state) + Error(_reason) -> Terminate(None) + } + } +} + +/// A client cancelling a stream. Enough of them in one window means Rapid Reset +/// (CVE-2023-44487), not ordinary cancelling, and the connection goes. +@internal +pub fn handle_client_reset(state: State, stream_id: Int) -> FrameResult { + let #(state, tripped) = record_reset(state) + + case + tripped, + dict.get(state.streams, stream_id), + stream_id > state.highest_client_stream_id_seen + { + True, _stream_lookup, _is_new_stream -> + Terminate(Some(frame.EnhanceYourCalm)) + False, Error(Nil), True -> Terminate(Some(frame.ProtocolError)) + False, Error(Nil), False -> Proceed(state) + False, Ok(entry), _is_new_stream -> + Proceed(reset_and_remove_stream(state, stream_id, entry)) + } +} + +fn record_reset(state: State) -> #(State, Bool) { + let now = monotonic_ms() + + let new_window = + state.reset_count == 0 + || now - state.reset_window_start > state.config.rapid_reset_window_ms + + let #(reset_window_start, reset_count) = case new_window { + True -> #(now, 1) + False -> #(state.reset_window_start, state.reset_count + 1) + } + + #( + State(..state, reset_window_start:, reset_count:), + reset_count > state.config.rapid_reset_threshold, + ) +} + +@external(erlang, "http2_ffi", "monotonic_ms") +fn monotonic_ms() -> Int + +fn remove_stream(state: State, stream_id: Int, entry: Stream) -> State { + let stream_pids = case entry.status { + Computing(pid) -> dict.delete(state.stream_pids, pid) + Flushing -> state.stream_pids + } + + State(..state, streams: dict.delete(state.streams, stream_id), stream_pids:) +} + +fn clear_stream_pid(state: State, pid: process.Pid) -> State { + State(..state, stream_pids: dict.delete(state.stream_pids, pid)) +} + +@internal +pub fn handle_headers( + state: State, + stream_id: Int, + end_stream: Bool, + end_headers: Bool, + payload: BitArray, +) -> FrameResult { + case + dict.get(state.streams, stream_id), + is_new_client_stream_id(state, stream_id) + { + Ok(entry), _is_new_stream if !entry.request_half_closed -> + start_header_assembly( + state, + stream_id, + end_stream, + end_headers, + payload, + True, + ) + Ok(_entry), _is_new_stream -> + RejectStream(state, stream_id, frame.StreamClosed) + Error(Nil), True -> { + State(..state, highest_client_stream_id_seen: stream_id) + |> start_header_assembly( + stream_id, + end_stream, + end_headers, + payload, + False, + ) + } + Error(Nil), False -> Terminate(Some(frame.ProtocolError)) + } +} + +fn start_header_assembly( + state: State, + stream_id: Int, + end_stream: Bool, + end_headers: Bool, + payload: BitArray, + trailers: Bool, +) -> FrameResult { + let assembly = + HeaderAssembly( + stream_id:, + end_stream:, + fragment_count: 1, + block: payload, + trailers:, + ) + + let oversized = + bit_array.byte_size(payload) > state.config.max_header_block_bytes + + case oversized, end_headers, trailers { + True, _end_headers, _trailers -> Terminate(Some(frame.EnhanceYourCalm)) + False, False, _trailers -> + Proceed(State(..state, header_assembly: Some(assembly))) + False, True, True -> complete_trailer_block(state, assembly) + False, True, False -> complete_header_block(state, assembly) + } +} + +// Clients only ever open odd-numbered streams and always in increasing order. +fn is_new_client_stream_id(state: State, stream_id: Int) -> Bool { + stream_id % 2 == 1 && stream_id > state.highest_client_stream_id_seen +} + +@internal +pub fn handle_data( + state: State, + stream_id: Int, + end_stream: Bool, + payload: BitArray, + size: Int, +) -> FrameResult { + case + stream_id == 0, + dict.get(state.streams, stream_id), + stream_id > state.highest_client_stream_id_seen + { + True, _lookup, _is_new_stream -> Terminate(Some(frame.ProtocolError)) + False, Error(Nil), True -> Terminate(Some(frame.ProtocolError)) + False, Error(Nil), False -> reject_closed_stream_data(state, size) + False, Ok(entry), _is_new_stream if entry.request_half_closed -> + reject_data_after_half_close(state, stream_id, size) + False, Ok(entry), _is_new_stream -> + apply_data(state, stream_id, entry, end_stream, payload, size) + } +} + +// The client sent this DATA frame before finding out the stream is closed on +// our end so it still counts against the connection window. Flow control has +// to be here even though the stream itself is already gone otherwise our view +// of the window drifts from the client's. +fn reject_closed_stream_data(state: State, size: Int) -> FrameResult { + case state.conn_recv_window - size < 0 { + True -> Terminate(Some(frame.FlowControlError)) + False -> Terminate(Some(frame.StreamClosed)) + } +} + +// Same here, still gotta debit the connection window before we can reject the +// stream. +fn reject_data_after_half_close( + state: State, + stream_id: Int, + size: Int, +) -> FrameResult { + let conn_recv_window = state.conn_recv_window - size + + case conn_recv_window < 0 { + True -> Terminate(Some(frame.FlowControlError)) + False -> + State(..state, conn_recv_window:) + |> RejectStream(stream_id, frame.StreamClosed) + } +} + +fn apply_data( + state: State, + stream_id: Int, + entry: Stream, + end_stream: Bool, + payload: BitArray, + size: Int, +) -> FrameResult { + let conn_recv_window = state.conn_recv_window - size + let recv_window = entry.recv_window - size + let state = State(..state, conn_recv_window:) + + case conn_recv_window < 0, recv_window < 0 { + True, _recv_window_negative -> Terminate(Some(frame.FlowControlError)) + False, True -> RejectStream(state, stream_id, frame.FlowControlError) + False, False -> + case content_length_violation(entry, payload, end_stream) { + True -> RejectStream(state, stream_id, frame.ProtocolError) + False -> { + let body_bytes_received = + entry.body_bytes_received + bit_array.byte_size(payload) + let #(state, conn_increment) = conn_recv_credit(state) + let #(entry, delivered_to_reader) = + Stream(..entry, recv_window:, body_bytes_received:) + |> deliver_data(payload, end_stream, []) + + let #(entry, stream_increment) = case delivered_to_reader { + True -> stream_recv_credit(entry, state.config) + False -> #(entry, 0) + } + + let streams = dict.insert(state.streams, stream_id, entry) + let state = State(..state, streams:) + + let out = + bytes_tree.new() + |> append_window_update(stream_id, stream_increment) + |> append_window_update(0, conn_increment) + + case stream_increment > 0 || conn_increment > 0 { + False -> Proceed(state) + True -> ProceedWithOutbound(state, out) + } + } + } + } +} + +fn content_length_violation( + entry: Stream, + payload: BitArray, + end_stream: Bool, +) -> Bool { + case entry.content_length { + None -> False + Some(expected) -> { + let total = entry.body_bytes_received + bit_array.byte_size(payload) + total > expected || { end_stream && total != expected } + } + } +} + +fn deliver_data( + entry: Stream, + payload: BitArray, + end_stream: Bool, + trailers: List(#(String, String)), +) -> #(Stream, Bool) { + let request_half_closed = entry.request_half_closed || end_stream + let entry = Stream(..entry, request_half_closed:) + + case entry.parked_reader { + Some(reply_to) -> { + case payload, end_stream { + <<>>, _half_closed -> process.send(reply_to, http2.DoneEvent(trailers)) + _payload, True -> + process.send(reply_to, http2.LastChunkEvent(payload, trailers)) + _payload, False -> process.send(reply_to, http2.ChunkEvent(payload)) + } + #(Stream(..entry, parked_reader: None), True) + } + None -> { + let recv_buffer = bytes_tree.append(entry.recv_buffer, payload) + #(Stream(..entry, recv_buffer:, trailers:), False) + } + } +} + +@internal +pub fn handle_continuation( + state: State, + assembly: HeaderAssembly, + frame: frame.Frame, +) -> FrameResult { + case frame { + frame.Continuation(stream_id, end_headers, payload) + if stream_id == assembly.stream_id + -> + case add_fragment(assembly, payload, state.config), end_headers { + Error(code), _end_headers -> Terminate(Some(code)) + Ok(updated), False -> + Proceed(State(..state, header_assembly: Some(updated))) + Ok(updated), True if updated.trailers -> + complete_trailer_block(state, updated) + Ok(updated), True -> complete_header_block(state, updated) + } + frame.Continuation(..) -> Terminate(Some(frame.ProtocolError)) + _frame -> Terminate(Some(frame.ProtocolError)) + } +} + +@internal +pub fn add_fragment( + assembly: HeaderAssembly, + fragment: BitArray, + config: http2.Config, +) -> Result(HeaderAssembly, frame.ErrorCode) { + let fragment_count = assembly.fragment_count + 1 + let block = <> + + case + fragment_count > config.max_continuation_frames + || bit_array.byte_size(block) > config.max_header_block_bytes + { + True -> Error(frame.EnhanceYourCalm) + False -> Ok(HeaderAssembly(..assembly, fragment_count:, block:)) + } +} + +fn decode_and_validate_header_block( + state: State, + assembly: HeaderAssembly, +) -> Result(#(State, List(#(BitArray, BitArray))), FrameResult) { + case alpacki.decode_header_block(assembly.block, state.hpack_decoder) { + Error(_decode_error) -> Error(Terminate(Some(frame.CompressionError))) + Ok(alpacki.DecodedHeaderBlock( + headers:, + decoded_size:, + dynamic_table:, + remaining:, + )) -> { + let header_list_size_exceeded = case state.config.max_header_list_size { + None -> False + Some(limit) -> decoded_size > limit + } + + case + remaining != <<>> + || alpacki.dynamic_max_size(dynamic_table) + > state.config.header_table_size, + header_list_size_exceeded + { + True, _header_list_size_exceeded -> + Error(Terminate(Some(frame.CompressionError))) + False, True -> Error(Terminate(Some(frame.EnhanceYourCalm))) + False, False -> { + let next_state = + State(..state, header_assembly: None, hpack_decoder: dynamic_table) + Ok(#(next_state, headers)) + } + } + } + } +} + +@internal +pub fn complete_header_block( + state: State, + assembly: HeaderAssembly, +) -> FrameResult { + case decode_and_validate_header_block(state, assembly) { + Error(result) -> result + Ok(#(next_state, headers)) -> { + let connection = + connection.Http2(http2.Connection( + connection: state.reply_subject, + stream_id: assembly.stream_id, + has_body: !assembly.end_stream, + pending: <<>>, + pending_trailers: None, + read: 0, + body_read_timeout: state.config.body_read_timeout, + peer: state.peer, + )) + + case build_request(headers, connection, state.patterns) { + Error(_error) -> + RejectStream(next_state, assembly.stream_id, frame.ProtocolError) + Ok(#(request, content_length)) -> + case assembly.end_stream, content_length { + True, Some(expected) if expected != 0 -> + RejectStream(next_state, assembly.stream_id, frame.ProtocolError) + _end_stream, _content_length -> + spawn_stream( + next_state, + assembly.stream_id, + request, + assembly.end_stream, + content_length, + ) + } + } + } + } +} + +fn complete_trailer_block( + state: State, + assembly: HeaderAssembly, +) -> FrameResult { + case decode_and_validate_header_block(state, assembly) { + Error(result) -> result + Ok(#(next_state, headers)) -> + case list.any(headers, is_pseudo_header), assembly.end_stream { + True, _end_stream -> + RejectStream(next_state, assembly.stream_id, frame.ProtocolError) + False, False -> + RejectStream(next_state, assembly.stream_id, frame.ProtocolError) + False, True -> + case decode_trailers(state.patterns, headers) { + Ok(trailers) -> + Proceed(finish_trailers(next_state, assembly.stream_id, trailers)) + Error(_error) -> + RejectStream(next_state, assembly.stream_id, frame.ProtocolError) + } + } + } +} + +fn is_pseudo_header(header: #(BitArray, BitArray)) -> Bool { + case header.0 { + <<58, _rest:bits>> -> True + _name -> False + } +} + +fn decode_trailers( + patterns: HeaderPatterns, + headers: List(#(BitArray, BitArray)), +) -> Result(List(#(String, String)), RequestError) { + let empty = + PseudoHeaders(method: None, scheme: None, authority: None, path: None) + |> HeaderAccumulated( + regular: dict.new(), + seen_regular: False, + content_length: None, + ) + + use acc <- result.try( + list.try_fold(headers, empty, fn(acc, header) { + let #(name, value) = header + add_regular(patterns, acc, name, value) + }), + ) + + Ok(dict.to_list(acc.regular)) +} + +fn finish_trailers( + state: State, + stream_id: Int, + trailers: List(#(String, String)), +) -> State { + case dict.get(state.streams, stream_id) { + Error(Nil) -> state + Ok(entry) -> { + let #(entry, _delivered_to_reader) = + deliver_data(entry, <<>>, True, trailers) + + State(..state, streams: dict.insert(state.streams, stream_id, entry)) + } + } +} + +fn spawn_stream( + state: State, + stream_id: Int, + request: Request(connection.Connection), + end_stream: Bool, + content_length: Option(Int), +) -> FrameResult { + let concurrent_streams_exceeded = case state.config.max_concurrent_streams { + None -> False + Some(limit) -> dict.size(state.streams) >= limit + } + + case state.draining || concurrent_streams_exceeded { + True -> RejectStream(state, stream_id, frame.RefusedStream) + False -> { + let pid = + stream.start(state.reply_subject, stream_id, request, state.handler) + Proceed(track_stream(state, stream_id, pid, end_stream, content_length)) + } + } +} + +fn track_stream( + state: State, + stream_id: Int, + pid: process.Pid, + end_stream: Bool, + content_length: Option(Int), +) -> State { + let entry = + Stream( + status: Computing(pid), + send_window: state.peer_settings.initial_window_size, + pending: PendingBytes(<<>>), + pending_end_stream: True, + write_ack: None, + recv_window: state.config.initial_window_size, + recv_buffer: bytes_tree.new(), + request_half_closed: end_stream, + parked_reader: None, + content_length:, + body_bytes_received: 0, + trailers: [], + ) + + State( + ..state, + streams: dict.insert(state.streams, stream_id, entry), + stream_pids: dict.insert(state.stream_pids, pid, stream_id), + ) +} + +@internal +pub type RequestError { + InvalidUtf8 + EmptyHeaderName + UppercaseHeaderName + MalformedHeaderBytes + PseudoHeaderAfterRegular + UnknownPseudoHeader + MissingPseudoHeader + DuplicatePseudoHeader + ConnectionSpecificHeader + InvalidMethod + InvalidScheme + InvalidAuthority + InvalidPath + InvalidContentLength +} + +type PseudoHeaders { + PseudoHeaders( + method: Option(http.Method), + scheme: Option(http.Scheme), + authority: Option(String), + path: Option(String), + ) +} + +type HeaderAccumulated { + HeaderAccumulated( + pseudo: PseudoHeaders, + regular: Dict(String, String), + seen_regular: Bool, + content_length: Option(Int), + ) +} + +@internal +pub fn build_request( + headers: List(#(BitArray, BitArray)), + body: body, + patterns: HeaderPatterns, +) -> Result(#(Request(body), Option(Int)), RequestError) { + let pseudo = + PseudoHeaders(method: None, scheme: None, authority: None, path: None) + + let empty = + HeaderAccumulated( + regular: dict.new(), + seen_regular: False, + content_length: None, + pseudo:, + ) + + use acc <- result.try( + list.try_fold(headers, empty, fn(acc, header) { + add_header(patterns, acc, header) + }), + ) + + case + acc.pseudo.method, + acc.pseudo.scheme, + acc.pseudo.authority, + acc.pseudo.path + { + Some(method), Some(scheme), Some(authority), Some(path) -> { + use #(host, port) <- result.try(split_authority(patterns, authority)) + let #(path, query) = case split_once(path, patterns.query) { + Ok(#(path, query)) -> #(path, Some(query)) + Error(Nil) -> #(path, None) + } + + Ok(#( + Request( + method:, + headers: dict.to_list(acc.regular), + body:, + scheme:, + host:, + port:, + path:, + query:, + ), + acc.content_length, + )) + } + _method, _scheme, _authority, _path -> Error(MissingPseudoHeader) + } +} + +fn add_header( + patterns: HeaderPatterns, + acc: HeaderAccumulated, + header: #(BitArray, BitArray), +) -> Result(HeaderAccumulated, RequestError) { + let #(name, value) = header + + case name { + <<>> -> Error(EmptyHeaderName) + <<":method":utf8>> -> + case acc.seen_regular, acc.pseudo.method { + True, _method -> Error(PseudoHeaderAfterRegular) + False, Some(_method) -> Error(DuplicatePseudoHeader) + False, None -> + case parse_method(patterns, value) { + Ok(method) -> { + let pseudo = PseudoHeaders(..acc.pseudo, method: Some(method)) + Ok(HeaderAccumulated(..acc, pseudo:)) + } + Error(error) -> Error(error) + } + } + <<":scheme":utf8>> -> + case acc.seen_regular, acc.pseudo.scheme { + True, _scheme -> Error(PseudoHeaderAfterRegular) + False, Some(_scheme) -> Error(DuplicatePseudoHeader) + False, None -> + case parse_scheme(value) { + Ok(scheme) -> { + let pseudo = PseudoHeaders(..acc.pseudo, scheme: Some(scheme)) + Ok(HeaderAccumulated(..acc, pseudo:)) + } + Error(error) -> Error(error) + } + } + <<":authority":utf8>> -> + case acc.seen_regular, acc.pseudo.authority { + True, _authority -> Error(PseudoHeaderAfterRegular) + False, Some(_authority) -> Error(DuplicatePseudoHeader) + False, None -> + case validate_header_value(patterns.forbidden, value) { + Ok(authority) -> { + let pseudo = + PseudoHeaders(..acc.pseudo, authority: Some(authority)) + Ok(HeaderAccumulated(..acc, pseudo:)) + } + Error(error) -> Error(error) + } + } + <<":path":utf8>> -> + case acc.seen_regular, acc.pseudo.path { + True, _path -> Error(PseudoHeaderAfterRegular) + False, Some(_path) -> Error(DuplicatePseudoHeader) + False, None -> + case validate_header_value(patterns.forbidden, value) { + Ok("") -> Error(InvalidPath) + Ok(path) -> { + let pseudo = PseudoHeaders(..acc.pseudo, path: Some(path)) + Ok(HeaderAccumulated(..acc, pseudo:)) + } + Error(error) -> Error(error) + } + } + <<58, _rest:bits>> -> + case acc.seen_regular { + True -> Error(PseudoHeaderAfterRegular) + False -> Error(UnknownPseudoHeader) + } + <<"connection":utf8>> + | <<"keep-alive":utf8>> + | <<"proxy-connection":utf8>> + | <<"transfer-encoding":utf8>> + | <<"upgrade":utf8>> -> Error(ConnectionSpecificHeader) + <<"content-length":utf8>> -> { + case acc.content_length { + Some(_prior) -> Error(InvalidContentLength) + None -> { + use value <- result.try(validate_header_value( + patterns.forbidden, + value, + )) + + case int.parse(value) { + Ok(n) if n >= 0 -> { + let regular = dict.insert(acc.regular, "content-length", value) + + HeaderAccumulated( + ..acc, + regular:, + seen_regular: True, + content_length: Some(n), + ) + |> Ok + } + _value -> Error(InvalidContentLength) + } + } + } + } + <<"te":utf8>> -> + case value { + <<"trailers":utf8>> -> { + let regular = dict.insert(acc.regular, "te", "trailers") + Ok(HeaderAccumulated(..acc, regular:)) + } + _name -> Error(ConnectionSpecificHeader) + } + _name -> add_regular(patterns, acc, name, value) + } +} + +fn parse_method( + patterns: HeaderPatterns, + value: BitArray, +) -> Result(http.Method, RequestError) { + case value { + <<"GET":utf8>> -> Ok(http.Get) + <<"POST":utf8>> -> Ok(http.Post) + <<"PUT":utf8>> -> Ok(http.Put) + <<"DELETE":utf8>> -> Ok(http.Delete) + <<"HEAD":utf8>> -> Ok(http.Head) + <<"OPTIONS":utf8>> -> Ok(http.Options) + <<"PATCH":utf8>> -> Ok(http.Patch) + <<"CONNECT":utf8>> -> Ok(http.Connect) + <<"TRACE":utf8>> -> Ok(http.Trace) + _method -> { + use method <- result.try(validate_header_value(patterns.forbidden, value)) + http.parse_method(method) |> result.replace_error(InvalidMethod) + } + } +} + +fn parse_scheme(value: BitArray) -> Result(http.Scheme, RequestError) { + case value { + <<"https":utf8>> -> Ok(http.Https) + <<"http":utf8>> -> Ok(http.Http) + _scheme -> Error(InvalidScheme) + } +} + +@external(erlang, "http2_ffi", "validate_header_name") +fn validate_header_name( + pattern: Pattern, + name: BitArray, +) -> Result(String, RequestError) + +@external(erlang, "http2_ffi", "validate_header_value") +fn validate_header_value( + pattern: Pattern, + value: BitArray, +) -> Result(String, RequestError) + +fn add_regular( + patterns: HeaderPatterns, + acc: HeaderAccumulated, + name: BitArray, + value: BitArray, +) -> Result(HeaderAccumulated, RequestError) { + use name <- result.try(validate_header_name(patterns.name, name)) + use value <- result.try(validate_header_value(patterns.forbidden, value)) + + let separator = case name { + "cookie" -> "; " + _existing -> ", " + } + + let regular = + dict_upsert( + name, + fn(prior) { prior <> separator <> value }, + value, + acc.regular, + ) + + Ok(HeaderAccumulated(..acc, regular:, seen_regular: True)) +} + +@external(erlang, "maps", "update_with") +fn dict_upsert( + key: String, + with: fn(String) -> String, + init: String, + map: Dict(String, String), +) -> Dict(String, String) + +fn split_authority( + patterns: HeaderPatterns, + authority: String, +) -> Result(#(String, Option(Int)), RequestError) { + case split_once(authority, patterns.colon) { + Ok(#(host, port_str)) -> + case int.parse(port_str) { + Ok(port) -> Ok(#(host, Some(port))) + Error(Nil) -> Error(InvalidAuthority) + } + Error(Nil) -> Ok(#(authority, None)) + } +} + +@external(erlang, "http2_ffi", "split_once") +fn split_once( + string: String, + on pattern: Pattern, +) -> Result(#(String, String), Nil) + +fn handle_client_settings( + state: State, + params: List(frame.Setting), + connection: glisten.Connection(connection.Message), +) -> FrameResult { + case send_frame(connection, frame.settings_ack) { + Error(_reason) -> Terminate(None) + Ok(Nil) -> { + let peer_settings = apply_settings(state.peer_settings, params) + + let delta = + peer_settings.initial_window_size + - state.peer_settings.initial_window_size + + let state = adjust_stream_windows(State(..state, peer_settings:), delta) + let state = case state.handshake { + AwaitingSettings -> { + cancel_timer(state.timer) + State(..state, handshake: Connected, timer: None) + } + Connected -> state + } + + case delta > 0 { + True -> flush_pending_streams(state, connection) + False -> Proceed(state) + } + } + } +} + +/// SETTINGS_INITIAL_WINDOW_SIZE applies retroactively, so every stream already +/// open gets the delta too. +@internal +pub fn adjust_stream_windows(state: State, delta: Int) -> State { + case delta { + 0 -> state + _delta -> { + let streams = + dict.map_values(state.streams, fn(_stream_id, entry) { + Stream(..entry, send_window: entry.send_window + delta) + }) + + State(..state, streams:) + } + } +} + +fn cancel_timer(timer: Option(process.Timer)) -> Nil { + case timer { + Some(timer) -> { + let _cancelled = process.cancel_timer(timer) + Nil + } + None -> Nil + } +} + +fn terminate( + state: State, + connection: glisten.Connection(connection.Message), + code: Option(frame.ErrorCode), +) -> Next { + case code { + Some(error_code) -> { + let _sent = send_frame(connection, frame.Goaway(0, 0, error_code, <<>>)) + Nil + } + None -> Nil + } + + stop_connection(state) +} + +fn send_frame( + connection: glisten.Connection(connection.Message), + frame: frame.Frame, +) -> Result(Nil, glisten.SocketReason) { + frame.encode(frame) + |> bytes_tree.from_bit_array + |> glisten.send(connection, _) +} + +fn begin_drain( + state: State, + connection: glisten.Connection(connection.Message), +) -> Next { + case state.draining { + True -> Continue(state) + False -> { + let timer = + process.send_after( + state.drain_subject, + state.config.drain_timeout_ms, + connection.Http2Drain, + ) + + let state = State(..state, draining: True, drain_timer: Some(timer)) + let goaway = + frame.Goaway( + 0, + state.highest_client_stream_id_seen, + frame.NoError, + <<>>, + ) + + case send_frame(connection, goaway) { + Ok(Nil) -> finish_or_continue(state) + Error(_reason) -> stop_connection(state) + } + } + } +} + +fn finish_or_continue(state: State) -> Next { + case state.draining && dict.size(state.streams) == 0 { + True -> { + cancel_timer(state.drain_timer) + Close + } + False -> Continue(state) + } +} + +fn handle_stream_reply( + state: State, + reply: http2.Reply(connection.Body), + connection: glisten.Connection(connection.Message), +) -> Next { + case reply { + http2.Respond(stream_id, response) -> + case dict.get(state.streams, stream_id) { + Error(Nil) -> Continue(state) + Ok(entry) -> respond(state, stream_id, entry, response, connection) + } + http2.ReadBody(stream_id, reply_to) -> + handle_read_body(state, stream_id, reply_to, connection) + http2.WriteHeaders(stream_id, ack, status, headers, reserved) -> + handle_write_headers( + state, + stream_id, + ack, + status, + headers, + reserved, + connection, + ) + http2.WriteData(stream_id, ack, chunk, end_stream) -> + handle_write_data(state, stream_id, ack, chunk, end_stream, connection) + } +} + +fn handle_read_body( + state: State, + stream_id: Int, + reply_to: process.Subject(http2.BodyEvent), + connection: glisten.Connection(connection.Message), +) -> Next { + case dict.get(state.streams, stream_id) { + Error(Nil) -> Continue(state) + Ok(entry) -> { + let buffered = bytes_tree.to_bit_array(entry.recv_buffer) + case buffered, entry.request_half_closed { + <<>>, True -> { + process.send(reply_to, http2.DoneEvent(entry.trailers)) + Continue(state) + } + <<>>, False -> { + let entry = Stream(..entry, parked_reader: Some(reply_to)) + let streams = dict.insert(state.streams, stream_id, entry) + Continue(State(..state, streams:)) + } + _buffered, True -> { + process.send(reply_to, http2.LastChunkEvent(buffered, entry.trailers)) + let drained_entry = Stream(..entry, recv_buffer: bytes_tree.new()) + let streams = dict.insert(state.streams, stream_id, drained_entry) + Continue(State(..state, streams:)) + } + _buffered, False -> { + process.send(reply_to, http2.ChunkEvent(buffered)) + let drained_entry = Stream(..entry, recv_buffer: bytes_tree.new()) + let #(drained_entry, stream_increment) = + stream_recv_credit(drained_entry, state.config) + + let streams = dict.insert(state.streams, stream_id, drained_entry) + let state = State(..state, streams:) + + case stream_increment > 0 { + False -> Continue(state) + True -> { + let out = + bytes_tree.new() + |> append_window_update(stream_id, stream_increment) + case glisten.send(connection, out) { + Ok(Nil) -> Continue(state) + Error(_reason) -> terminate(state, connection, None) + } + } + } + } + } + } + } +} + +/// Tops a stream's receive window back up once it's fallen far enough to be +/// worth a frame. Beats trickling out an update per chunk read. +@internal +pub fn stream_recv_credit( + entry: Stream, + config: http2.Config, +) -> #(Stream, Int) { + case entry.recv_window <= config.recv_window_low_water_mark { + True -> { + let increment = config.recv_window_high_water_mark - entry.recv_window + + #( + Stream(..entry, recv_window: config.recv_window_high_water_mark), + increment, + ) + } + False -> #(entry, 0) + } +} + +/// The same for the connection's own window which every stream draws from. +@internal +pub fn conn_recv_credit(state: State) -> #(State, Int) { + case state.conn_recv_window <= state.config.recv_window_low_water_mark { + True -> { + let increment = + state.config.recv_window_high_water_mark - state.conn_recv_window + + #( + State( + ..state, + conn_recv_window: state.config.recv_window_high_water_mark, + ), + increment, + ) + } + False -> #(state, 0) + } +} + +fn append_window_update( + acc: bytes_tree.BytesTree, + stream_id: Int, + increment: Int, +) -> bytes_tree.BytesTree { + case increment > 0 { + True -> + frame.encode(frame.WindowUpdate(stream_id, increment)) + |> bytes_tree.append(acc, _) + False -> acc + } +} + +fn build_response_headers( + headers: List(#(String, String)), + pattern: Pattern, +) -> List(alpacki.HeaderField) { + list.fold(headers, [], fn(fields, header) { + let #(name, value) = header + case name { + "connection" + | "keep-alive" + | "proxy-connection" + | "transfer-encoding" + | "upgrade" + | "date" + | "content-length" + | "" -> fields + _name -> + case + has_forbidden_header_bytes(pattern, name) + || has_forbidden_header_bytes(pattern, value) + { + True -> fields + False -> [ + alpacki.HeaderField( + <>, + <>, + alpacki.WithoutIndexing, + ), + ..fields + ] + } + } + }) +} + +@external(erlang, "http2_ffi", "has_forbidden_header_bytes") +fn has_forbidden_header_bytes(pattern: Pattern, value: String) -> Bool + +fn response_body_size(body: connection.Body) -> Int { + case body { + connection.Bytes(tree) -> bytes_tree.byte_size(tree) + connection.Text(text) -> byte_size(text) + connection.Empty -> 0 + connection.File(connection.OpenFile(length:, ..)) + | connection.File(connection.PendingFile(length:, ..)) -> length + // A streamed body never reaches here as its stream process writes the + // response itself + connection.Streaming(_metadata) + | connection.Sse(_metadata) + | connection.Websocket(_metadata) -> + panic as "a streamed body is written by its own stream process" + } +} + +@external(erlang, "erlang", "byte_size") +fn byte_size(text: String) -> Int + +fn open_pending( + body: connection.Body, + file_read_threshold: Int, +) -> Result(Pending, file.FileError) { + case body { + connection.Bytes(tree) -> Ok(PendingBytes(bytes_tree.to_bit_array(tree))) + connection.Text(text) -> Ok(PendingBytes(bit_array.from_string(text))) + connection.Empty -> Ok(PendingBytes(<<>>)) + connection.File(connection.OpenFile(handle:, offset:, length:)) -> + Ok(PendingFile(handle, offset, length)) + // Small enough to answer from memory. + connection.File(connection.PendingFile(path:, offset:, length:)) + if length <= file_read_threshold + -> file.read_range(path, offset, length) |> result.map(PendingBytes) + connection.File(connection.PendingFile(path:, offset:, length:)) -> + file.open(path) |> result.map(PendingFile(_, offset, length)) + connection.Streaming(_metadata) + | connection.Sse(_metadata) + | connection.Websocket(_metadata) -> + panic as "a streamed body is written by its own stream process" + } +} + +fn respond( + state: State, + stream_id: Int, + entry: Stream, + response: Response(connection.Body), + connection: glisten.Connection(connection.Message), +) -> Next { + case open_pending(response.body, state.config.file_read_threshold) { + Ok(pending) -> + send_response( + state, + stream_id, + entry, + response.status, + response.headers, + response_body_size(response.body), + pending, + connection, + ) + Error(_error) -> + send_response( + state, + stream_id, + entry, + 500, + [], + 0, + PendingBytes(<<>>), + connection, + ) + } +} + +fn send_response( + state: State, + stream_id: Int, + entry: Stream, + status: Int, + headers: List(#(String, String)), + body_size: Int, + pending: Pending, + connection: glisten.Connection(connection.Message), +) -> Next { + let fields = build_response_headers(headers, state.patterns.forbidden) + + let content_length_fields = case status { + status if status == 204 || { status >= 100 && status < 200 } -> [] + _status -> [ + alpacki.HeaderField( + <<"content-length":utf8>>, + <>, + alpacki.WithoutIndexing, + ), + ] + } + + let header_fields = [ + alpacki.HeaderField( + <<":status":utf8>>, + <>, + alpacki.WithoutIndexing, + ), + alpacki.HeaderField(<<"date":utf8>>, clock.get(), alpacki.WithoutIndexing), + ..list.append(content_length_fields, list.reverse(fields)) + ] + + let #(block, hpack_encoder) = + alpacki.encode_header_block(header_fields, state.hpack_encoder, True) + + let state = State(..state, hpack_encoder:) + + let has_body = body_size != 0 + let out = + bytes_tree.new() + |> append_header_frames( + stream_id, + !has_body, + block, + state.peer_settings.max_frame_size, + ) + + case has_body { + False -> { + let state = State(..state, streams: dict.delete(state.streams, stream_id)) + case glisten.send(connection, out) { + Ok(Nil) -> finish_or_continue(state) + Error(_reason) -> terminate(state, connection, None) + } + } + True -> { + let pending_entry = Stream(..entry, status: Flushing, pending:) + + flush_stream(state, stream_id, pending_entry, out, connection) + |> resolve_frame_result(state, connection) + } + } +} + +/// A stream sets these itself, so the handler's copies get dropped. HTTP/1 +/// reserves `transfer-encoding` too, which HTTP/2 has no use for and must +/// never send. +fn drop_sse_headers( + headers: List(#(String, String)), +) -> List(#(String, String)) { + use #(name, _value) <- list.filter(headers) + name != "content-type" && name != "cache-control" +} + +fn sse_fields() -> List(alpacki.HeaderField) { + [ + alpacki.HeaderField( + <<"content-type":utf8>>, + <<"text/event-stream":utf8>>, + alpacki.WithoutIndexing, + ), + alpacki.HeaderField( + <<"cache-control":utf8>>, + <<"no-cache":utf8>>, + alpacki.WithoutIndexing, + ), + ] +} + +fn handle_write_headers( + state: State, + stream_id: Int, + ack: process.Subject(http2.WriteAck), + status: Int, + headers: List(#(String, String)), + reserved: http2.Reserved, + connection: glisten.Connection(connection.Message), +) -> Next { + case dict.get(state.streams, stream_id) { + Error(Nil) -> Continue(state) + Ok(_entry) -> { + let headers = case reserved { + http2.Nothing -> headers + http2.SseHeaders -> drop_sse_headers(headers) + } + let fields = build_response_headers(headers, state.patterns.forbidden) + + let reserved_fields = case reserved { + http2.Nothing -> [] + http2.SseHeaders -> sse_fields() + } + + let header_fields = [ + alpacki.HeaderField( + <<":status":utf8>>, + <>, + alpacki.WithoutIndexing, + ), + alpacki.HeaderField( + <<"date":utf8>>, + clock.get(), + alpacki.WithoutIndexing, + ), + ..list.append(reserved_fields, list.reverse(fields)) + ] + + let #(block, hpack_encoder) = + alpacki.encode_header_block(header_fields, state.hpack_encoder, True) + + let state = State(..state, hpack_encoder:) + + let out = + bytes_tree.new() + |> append_header_frames( + stream_id, + False, + block, + state.peer_settings.max_frame_size, + ) + + case glisten.send(connection, out) { + Ok(Nil) -> { + process.send(ack, http2.WriteAck) + Continue(state) + } + Error(_reason) -> terminate(state, connection, None) + } + } + } +} + +fn handle_write_data( + state: State, + stream_id: Int, + reply_to: process.Subject(http2.WriteAck), + chunk: BitArray, + end_stream: Bool, + connection: glisten.Connection(connection.Message), +) -> Next { + case dict.get(state.streams, stream_id) { + Error(Nil) -> Continue(state) + Ok(entry) -> { + let entry = + Stream( + ..entry, + pending: PendingBytes(chunk), + pending_end_stream: end_stream, + write_ack: Some(reply_to), + ) + + flush_stream(state, stream_id, entry, bytes_tree.new(), connection) + |> resolve_frame_result(state, connection) + } + } +} + +fn resolve_frame_result( + result: FrameResult, + state: State, + connection: glisten.Connection(connection.Message), +) -> Next { + case result { + Proceed(state) -> finish_or_continue(state) + ProceedWithOutbound(state, out) -> + case glisten.send(connection, out) { + Ok(Nil) -> finish_or_continue(state) + Error(_reason) -> terminate(state, connection, None) + } + RejectStream(state, stream_id, code) -> + reject_stream(connection, state, stream_id, code) + Terminate(code) -> terminate(state, connection, code) + } +} + +/// Frames a header block, splitting into CONTINUATION frames when it's longer +/// than the peer takes in one. +@internal +pub fn append_header_frames( + acc: bytes_tree.BytesTree, + stream_id: Int, + end_stream: Bool, + block: BitArray, + max_frame_size: Int, +) -> bytes_tree.BytesTree { + case block { + <> if remaining != <<>> -> + frame.encode(frame.Headers(stream_id, end_stream, False, chunk)) + |> bytes_tree.append(acc, _) + |> append_continuation_frames(stream_id, remaining, max_frame_size) + _block -> + frame.encode(frame.Headers(stream_id, end_stream, True, block)) + |> bytes_tree.append(acc, _) + } +} + +fn append_continuation_frames( + acc: bytes_tree.BytesTree, + stream_id: Int, + block: BitArray, + max_frame_size: Int, +) -> bytes_tree.BytesTree { + case block { + <> if remaining != <<>> -> + frame.encode(frame.Continuation(stream_id, False, chunk)) + |> bytes_tree.append(acc, _) + |> append_continuation_frames(stream_id, remaining, max_frame_size) + _block -> + frame.encode(frame.Continuation(stream_id, True, block)) + |> bytes_tree.append(acc, _) + } +} + +@internal +pub type FlushOutcome { + FlushAccumulated(state: State, out: bytes_tree.BytesTree, wrote: Bool) + FlushFileChunk( + state: State, + out: bytes_tree.BytesTree, + stream_id: Int, + entry: Stream, + chunk_size: Int, + end_stream: Bool, + ) +} + +@internal +pub fn do_flush_stream( + state: State, + stream_id: Int, + entry: Stream, + out: bytes_tree.BytesTree, + wrote: Bool, +) -> FlushOutcome { + case entry.pending { + PendingBytes(<<>>) -> { + let out = + frame.encode(frame.Data(stream_id, entry.pending_end_stream, <<>>, 0)) + |> bytes_tree.append(out, _) + + finish_pending(state, stream_id, entry, out, True) + } + PendingFile(descriptor, _offset, 0) -> { + file.close(descriptor) + + State(..state, streams: dict.delete(state.streams, stream_id)) + |> FlushAccumulated(out, wrote) + } + PendingBytes(pending) -> { + let allowed = + int.min(state.conn_send_window, entry.send_window) + |> int.min(state.peer_settings.max_frame_size) + + case allowed <= 0 { + True -> { + State(..state, streams: dict.insert(state.streams, stream_id, entry)) + |> FlushAccumulated(out, wrote) + } + False -> + case pending { + <> if remaining != <<>> -> { + let out = + frame.encode_data_header(stream_id, False, allowed) + |> bytes_tree.append(out, _) + |> bytes_tree.append(chunk) + + let entry = + Stream( + ..entry, + send_window: entry.send_window - allowed, + pending: PendingBytes(remaining), + ) + + let state = + State( + ..state, + conn_send_window: state.conn_send_window - allowed, + ) + + do_flush_stream(state, stream_id, entry, out, True) + } + _pending -> { + let sent = bit_array.byte_size(pending) + let out = + frame.encode_data_header( + stream_id, + entry.pending_end_stream, + sent, + ) + |> bytes_tree.append(out, _) + |> bytes_tree.append(pending) + + let entry = + Stream( + ..entry, + send_window: entry.send_window - sent, + pending: PendingBytes(<<>>), + ) + + let state = + State(..state, conn_send_window: state.conn_send_window - sent) + + finish_pending(state, stream_id, entry, out, True) + } + } + } + } + PendingFile(_descriptor, _offset, remaining) -> { + let allowed = + int.min(state.conn_send_window, entry.send_window) + |> int.min(state.peer_settings.max_frame_size) + + case allowed <= 0 { + True -> { + State(..state, streams: dict.insert(state.streams, stream_id, entry)) + |> FlushAccumulated(out, wrote) + } + False -> { + let chunk_size = int.min(allowed, remaining) + let end_stream = chunk_size == remaining + let out = + frame.encode_data_header(stream_id, end_stream, chunk_size) + |> bytes_tree.append(out, _) + + FlushFileChunk(state, out, stream_id, entry, chunk_size, end_stream) + } + } + } + } +} + +fn finish_pending( + state: State, + stream_id: Int, + entry: Stream, + out: bytes_tree.BytesTree, + wrote: Bool, +) -> FlushOutcome { + case entry.write_ack { + Some(reply_to) -> process.send(reply_to, http2.WriteAck) + None -> Nil + } + + case entry.pending_end_stream { + True -> + State(..state, streams: dict.delete(state.streams, stream_id)) + |> FlushAccumulated(out, wrote) + False -> { + let entry = Stream(..entry, write_ack: None) + State(..state, streams: dict.insert(state.streams, stream_id, entry)) + |> FlushAccumulated(out, wrote) + } + } +} + +fn send_if_any( + connection: glisten.Connection(connection.Message), + state: State, + out: bytes_tree.BytesTree, + wrote: Bool, +) -> FrameResult { + case wrote { + False -> Proceed(state) + True -> + case glisten.send(connection, out) { + Ok(Nil) -> Proceed(state) + Error(_reason) -> Terminate(None) + } + } +} + +fn flush_many( + state: State, + streams: List(#(Int, Stream)), + out: bytes_tree.BytesTree, + wrote: Bool, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + case streams { + [] -> send_if_any(connection, state, out, wrote) + [#(stream_id, entry), ..remaining] -> + case do_flush_stream(state, stream_id, entry, out, wrote) { + FlushAccumulated(state, out, wrote) -> + flush_many(state, remaining, out, wrote, connection) + FlushFileChunk(state, out, stream_id, entry, chunk_size, end_stream) -> + send_file_chunk( + state, + stream_id, + entry, + out, + chunk_size, + end_stream, + remaining, + connection, + ) + } + } +} + +fn send_file_chunk( + state: State, + stream_id: Int, + entry: Stream, + out: bytes_tree.BytesTree, + chunk_size: Int, + end_stream: Bool, + rest: List(#(Int, Stream)), + connection: glisten.Connection(connection.Message), +) -> FrameResult { + let assert PendingFile(descriptor, offset, remaining) = entry.pending + + case glisten.send(connection, out) { + Error(_reason) -> { + file.close(descriptor) + Terminate(None) + } + Ok(Nil) -> + case + file.send_chunk( + connection.transport, + connection.socket, + descriptor, + offset, + chunk_size, + ) + { + Error(_reason) -> { + file.close(descriptor) + Terminate(None) + } + Ok(Nil) -> { + let state = + State( + ..state, + conn_send_window: state.conn_send_window - chunk_size, + ) + + case end_stream { + True -> { + file.close(descriptor) + + State(..state, streams: dict.delete(state.streams, stream_id)) + |> flush_many(rest, bytes_tree.new(), False, connection) + } + False -> { + let pending = + PendingFile( + descriptor, + offset + chunk_size, + remaining - chunk_size, + ) + let entry = + Stream( + ..entry, + send_window: entry.send_window - chunk_size, + pending:, + ) + + flush_many( + state, + [#(stream_id, entry), ..rest], + bytes_tree.new(), + False, + connection, + ) + } + } + } + } + } +} + +/// Writes what a stream has pending as far as its window and the connection's +/// allow. +@internal +pub fn flush_stream( + state: State, + stream_id: Int, + entry: Stream, + out: bytes_tree.BytesTree, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + flush_many(state, [#(stream_id, entry)], out, False, connection) +} + +fn has_pending(entry: Stream) -> Bool { + case entry.pending { + PendingBytes(<<>>) -> False + PendingFile(_descriptor, _offset, 0) -> False + _pending -> True + } +} + +fn handle_window_update( + state: State, + stream_id: Int, + increment: Int, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + case increment > 0, stream_id { + False, 0 -> Terminate(Some(frame.ProtocolError)) + False, _stream_id -> RejectStream(state, stream_id, frame.ProtocolError) + True, 0 -> connection_window_update(state, increment, connection) + True, _stream_id -> + stream_window_update(state, stream_id, increment, connection) + } +} + +fn connection_window_update( + state: State, + increment: Int, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + let new_window = state.conn_send_window + increment + + case new_window > max_window_size { + True -> Terminate(Some(frame.FlowControlError)) + False -> + flush_pending_streams( + State(..state, conn_send_window: new_window), + connection, + ) + } +} + +fn stream_window_update( + state: State, + stream_id: Int, + increment: Int, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + case + dict.get(state.streams, stream_id), + stream_id > state.highest_client_stream_id_seen + { + Error(Nil), True -> Terminate(Some(frame.ProtocolError)) + Error(Nil), False -> Proceed(state) + Ok(entry), _is_new_stream -> { + let new_window = entry.send_window + increment + + case new_window > max_window_size { + True -> RejectStream(state, stream_id, frame.FlowControlError) + False -> { + let entry = Stream(..entry, send_window: new_window) + + case entry.status, has_pending(entry) { + Computing(_pid), False -> { + let streams = dict.insert(state.streams, stream_id, entry) + Proceed(State(..state, streams:)) + } + Computing(_pid), True | Flushing, _has_pending -> + flush_stream( + state, + stream_id, + entry, + bytes_tree.new(), + connection, + ) + } + } + } + } + } +} + +fn flush_pending_streams( + state: State, + connection: glisten.Connection(connection.Message), +) -> FrameResult { + let pending = + dict.filter(state.streams, fn(_stream_id, entry) { has_pending(entry) }) + |> dict.to_list + + flush_many(state, pending, bytes_tree.new(), False, connection) +} + +fn handle_stream_exit( + state: State, + exit: process.ExitMessage, + connection: glisten.Connection(connection.Message), +) -> Next { + case dict.get(state.stream_pids, exit.pid) { + // Not a stream this connection started so it is the parent asking it to + // shut down. + Error(Nil) -> begin_drain(state, connection) + Ok(stream_id) -> + case dict.get(state.streams, stream_id) { + // Still computing when its process ended means the handler returned + // without a response being completed. + Ok(Stream(status: Computing(pid), ..) as entry) if pid == exit.pid -> + remove_stream(state, stream_id, entry) + |> reject_stream(connection, _, stream_id, frame.InternalError) + _entry -> finish_or_continue(clear_stream_pid(state, exit.pid)) + } + } +} + +@internal +pub fn test_state() -> State { + let config = http2.default_config() + + State( + buffer: <<>>, + handshake: Connected, + peer_settings: default_peer_settings, + timer: None, + hpack_decoder: alpacki.new_dynamic(config.header_table_size), + hpack_encoder: alpacki.new_dynamic(config.header_table_size), + header_assembly: None, + reply_subject: process.new_subject(), + handler: fn(_request) { + response.new(200) |> response.set_body(connection.Empty) + }, + streams: dict.new(), + conn_send_window: default_send_window, + conn_recv_window: default_send_window, + stream_pids: dict.new(), + reset_window_start: 0, + reset_count: 0, + highest_client_stream_id_seen: 0, + draining: False, + drain_subject: process.new_subject(), + drain_timer: None, + config:, + settings_frame: build_settings_frame(config), + patterns: header_patterns(), + peer: Error(Nil), + ) +} diff --git a/src/ewe/internal/http2/body.gleam b/src/ewe/internal/http2/body.gleam new file mode 100644 index 0000000..a3fa2e4 --- /dev/null +++ b/src/ewe/internal/http2/body.gleam @@ -0,0 +1,121 @@ +import ewe/internal/http2/connection as http2 +import gleam/bit_array +import gleam/bytes_tree +import gleam/erlang/process +import gleam/erlang/reference +import gleam/option + +pub type BodyError { + BodyTooLarge + InvalidBody +} + +pub fn read_body( + connection: http2.Connection(body), + limit: Int, +) -> Result(#(BitArray, List(#(String, String))), BodyError) { + case connection.has_body { + False -> Ok(#(<<>>, [])) + True -> read_all(connection, limit, bytes_tree.new()) + } +} + +fn read_all( + connection: http2.Connection(body), + limit: Int, + acc: bytes_tree.BytesTree, +) -> Result(#(BitArray, List(#(String, String))), BodyError) { + case next_chunk(connection) { + Error(error) -> Error(error) + Ok(Done(trailers)) -> Ok(#(bytes_tree.to_bit_array(acc), trailers)) + Ok(Chunk(data, connection)) -> + case connection.read > limit { + True -> Error(BodyTooLarge) + False -> read_all(connection, limit, bytes_tree.append(acc, data)) + } + } +} + +/// The result of one `read_body_chunk` call. +pub type ReadEvent(body) { + Chunk(data: BitArray, connection: http2.Connection(body)) + Done(trailers: List(#(String, String))) +} + +pub fn read_body_chunk( + connection: http2.Connection(body), + max_chunk_bytes max_chunk_bytes: Int, + limit limit: Int, +) -> Result(ReadEvent(body), BodyError) { + case next_chunk(connection) { + Error(error) -> Error(error) + Ok(Done(trailers)) -> Ok(Done(trailers)) + Ok(Chunk(data, connection)) -> + case connection.read > limit { + True -> Error(BodyTooLarge) + False -> Ok(split(connection, data, max_chunk_bytes)) + } + } +} + +/// Leftovers past `max_chunk_bytes` wait on the connection for the next call. +fn split( + connection: http2.Connection(body), + data: BitArray, + max_chunk_bytes: Int, +) -> ReadEvent(body) { + case data { + <> -> + Chunk(chunk, http2.Connection(..connection, pending:)) + _data -> Chunk(data, http2.Connection(..connection, pending: <<>>)) + } +} + +fn next_chunk( + connection: http2.Connection(body), +) -> Result(ReadEvent(body), BodyError) { + case connection.has_body, connection.pending, connection.pending_trailers { + False, _pending, _trailers -> Ok(Done([])) + True, <<>>, option.Some(trailers) -> Ok(Done(trailers)) + True, <<>>, option.None -> pull(connection) + // Owed from the last split. Hand these back before asking for more. + True, pending, _trailers -> + Ok(Chunk(pending, http2.Connection(..connection, pending: <<>>))) + } +} + +fn pull( + connection: http2.Connection(body), +) -> Result(ReadEvent(body), BodyError) { + let tag = reference.new() + let reply_to = process.unsafely_create_subject(process.self(), http2.tag(tag)) + + process.send( + connection.connection, + http2.ReadBody(connection.stream_id, reply_to), + ) + + case http2.receive_reply_within(tag, connection.body_read_timeout) { + Error(_interrupted) -> Error(InvalidBody) + Ok(http2.DoneEvent(trailers)) -> Ok(Done(trailers)) + Ok(http2.ChunkEvent(data)) -> Ok(Chunk(data, advance(connection, data))) + Ok(http2.LastChunkEvent(data, trailers)) -> + Ok(Chunk( + data, + http2.Connection( + ..advance(connection, data), + pending_trailers: option.Some(trailers), + ), + )) + } +} + +fn advance( + connection: http2.Connection(body), + data: BitArray, +) -> http2.Connection(body) { + http2.Connection( + ..connection, + read: connection.read + bit_array.byte_size(data), + ) +} diff --git a/src/ewe/internal/http2/connection.gleam b/src/ewe/internal/http2/connection.gleam new file mode 100644 index 0000000..0c129a7 --- /dev/null +++ b/src/ewe/internal/http2/connection.gleam @@ -0,0 +1,146 @@ +import gleam/dynamic +import gleam/erlang/process +import gleam/erlang/reference +import gleam/http/response +import gleam/option +import glisten/socket + +/// The limits and timeouts an HTTP/2 connection is held to. +pub type Config { + Config( + max_concurrent_streams: option.Option(Int), + initial_window_size: Int, + max_frame_size: Int, + max_header_list_size: option.Option(Int), + header_table_size: Int, + max_continuation_frames: Int, + max_header_block_bytes: Int, + rapid_reset_window_ms: Int, + rapid_reset_threshold: Int, + handshake_timeout_ms: Int, + drain_timeout_ms: Int, + recv_window_low_water_mark: Int, + recv_window_high_water_mark: Int, + file_read_threshold: Int, + body_read_timeout: Int, + ) +} + +pub fn default_config() -> Config { + Config( + max_concurrent_streams: option.None, + initial_window_size: 2_097_152, + max_frame_size: 16_384, + max_header_list_size: option.Some(32_768), + header_table_size: 4096, + max_continuation_frames: 100, + max_header_block_bytes: 65_536, + rapid_reset_window_ms: 10_000, + rapid_reset_threshold: 100, + handshake_timeout_ms: 10_000, + drain_timeout_ms: 4000, + recv_window_low_water_mark: 262_144, + recv_window_high_water_mark: 2_097_152, + file_read_threshold: 1_048_576, + body_read_timeout: 10_000, + ) +} + +/// What a handler holds for one stream. The handler gets its own process, not +/// the connection's. Reading the body and writing the response are messages, +/// not socket writes. +/// +/// The body type is a parameter to break the import cycle with the module that +/// defines it. +pub type Connection(body) { + Connection( + connection: process.Subject(Reply(body)), + stream_id: Int, + has_body: Bool, + /// Body bytes handed over but not yet returned to the caller. + pending: BitArray, + pending_trailers: option.Option(List(#(String, String))), + /// Body bytes read so far. `read_body_chunk` caps its limit against this. + read: Int, + body_read_timeout: Int, + /// Resolved once for the connection. A stream process has no socket to + /// ask. + peer: Result(socket.SockName, Nil), + ) +} + +/// What a stream process asks of the connection process. +pub type Reply(body) { + Respond(stream_id: Int, response: response.Response(body)) + ReadBody(stream_id: Int, reply_to: process.Subject(BodyEvent)) + WriteHeaders( + stream_id: Int, + ack: process.Subject(WriteAck), + status: Int, + headers: List(#(String, String)), + reserved: Reserved, + ) + WriteData( + stream_id: Int, + ack: process.Subject(WriteAck), + chunk: BitArray, + end_stream: Bool, + ) +} + +/// Which headers the connection sets itself for a streamed body. They get +/// dropped from the handler's list so nothing goes out twice. +pub type Reserved { + Nothing + SseHeaders +} + +pub type BodyEvent { + ChunkEvent(BitArray) + LastChunkEvent(BitArray, trailers: List(#(String, String))) + DoneEvent(trailers: List(#(String, String))) +} + +pub type WriteAck { + WriteAck +} + +/// A handle for writing a response a frame at a time. The connection tags its +/// acks with the reference so a stream waits on it directly, no selector. +pub type ResponseWriter(body) { + ResponseWriter( + connection: process.Subject(Reply(body)), + stream_id: Int, + ack: process.Subject(WriteAck), + ack_ref: reference.Reference, + ) +} + +pub type SseConnection(body) { + SseConnection(writer: ResponseWriter(body)) +} + +/// Why a stream process stopped waiting. Only a body read times out. It is the +/// one wait that hangs on the client. +pub type Interrupted { + StreamReset + ConnectionClosed + TimedOut +} + +/// Waits for the connection to answer. Gives up if the stream resets or the +/// connection goes. +@external(erlang, "http2_ffi", "recv_or_exit") +pub fn receive_reply(tag: reference.Reference) -> Result(message, Interrupted) + +@external(erlang, "http2_ffi", "recv_or_exit") +pub fn receive_reply_within( + tag: reference.Reference, + timeout: Int, +) -> Result(message, Interrupted) + +/// `unsafely_create_subject` wants the tag the messages carry and the +/// connection tags replies with a reference, so the two meet as `Dynamic`. +/// Waiting on the reference directly saves building a selector per chunk. +@external(erlang, "ewe_ffi", "identity") +pub fn tag(reference: reference.Reference) -> dynamic.Dynamic diff --git a/src/ewe/internal/http2/frame.gleam b/src/ewe/internal/http2/frame.gleam new file mode 100644 index 0000000..765f60b --- /dev/null +++ b/src/ewe/internal/http2/frame.gleam @@ -0,0 +1,448 @@ +import gleam/bit_array +import gleam/list +import gleam/result + +pub type Frame { + Data( + stream_id: Int, + end_stream: Bool, + payload: BitArray, + flow_control_size: Int, + ) + Headers( + stream_id: Int, + end_stream: Bool, + end_headers: Bool, + payload: BitArray, + ) + Priority(stream_id: Int) + RstStream(stream_id: Int, error_code: ErrorCode) + Settings(stream_id: Int, ack: Bool, params: List(Setting)) + PushPromise(stream_id: Int, payload: BitArray) + Ping(stream_id: Int, ack: Bool, opaque_data: BitArray) + Goaway( + stream_id: Int, + last_stream_id: Int, + error_code: ErrorCode, + debug_data: BitArray, + ) + WindowUpdate(stream_id: Int, increment: Int) + Continuation(stream_id: Int, end_headers: Bool, payload: BitArray) + Unknown(stream_id: Int, frame_type: Int, payload: BitArray) +} + +pub const settings_ack = Settings(0, True, []) + +pub type Setting { + HeaderTableSize(Int) + EnablePush(Bool) + MaxConcurrentStreams(Int) + InitialWindowSize(Int) + MaxFrameSize(Int) + MaxHeaderListSize(Int) + UnknownSetting(Int, Int) +} + +pub type ErrorCode { + NoError + ProtocolError + InternalError + FlowControlError + SettingsTimeout + StreamClosed + FrameSizeError + RefusedStream + Cancel + CompressionError + ConnectError + EnhanceYourCalm + InadequateSecurity + Http11Required + UnknownErrorCode(Int) +} + +pub type FrameError { + Incomplete + Violation(ErrorCode) +} + +pub fn decode( + data: BitArray, + max_frame_size: Int, +) -> Result(#(Frame, BitArray), FrameError) { + case data { + << + length:size(24), + type_:8, + _unused_flags:2, + priority:1, + _unused_flag_high:1, + padded:1, + end_headers:1, + _unused_flag_low:1, + end_stream_or_ack:1, + _reserved_bit:1, + stream_id:31, + remaining:bits, + >> -> + case length > max_frame_size { + True -> Error(Violation(FrameSizeError)) + False -> + case remaining { + <> -> { + use frame <- result.try(decode_payload( + type_, + stream_id, + end_stream_or_ack == 1, + end_headers == 1, + padded == 1, + priority == 1, + payload, + )) + Ok(#(frame, remaining)) + } + _remaining -> Error(Incomplete) + } + } + _data -> Error(Incomplete) + } +} + +fn decode_payload( + type_: Int, + stream_id: Int, + end_stream_or_ack: Bool, + end_headers: Bool, + padded: Bool, + priority: Bool, + payload: BitArray, +) -> Result(Frame, FrameError) { + case type_ { + 0x0 -> { + use content <- result.try(strip_padding(padded, payload)) + Ok(Data( + stream_id:, + end_stream: end_stream_or_ack, + payload: content, + flow_control_size: bit_array.byte_size(payload), + )) + } + 0x1 -> { + use content <- result.try(strip_padding(padded, payload)) + use header_block <- result.try(strip_priority( + priority, + stream_id, + content, + )) + Ok(Headers( + stream_id:, + end_stream: end_stream_or_ack, + end_headers:, + payload: header_block, + )) + } + 0x2 -> + case payload { + <<_exclusive:1, stream_dependency:31, _weight:8>> -> + // Stream 0 can't be prioritized and a stream depending on itself + // is a cycle of one which both are protocol errors. + case stream_id == 0, stream_dependency == stream_id { + True, _self_dependent | _on_stream_zero, True -> + Error(Violation(ProtocolError)) + False, False -> Ok(Priority(stream_id:)) + } + _payload -> Error(Violation(FrameSizeError)) + } + 0x3 -> + case stream_id == 0, payload { + True, _payload -> Error(Violation(ProtocolError)) + False, <> -> + Ok(RstStream(stream_id:, error_code: decode_error_code(error_code))) + False, _payload -> Error(Violation(FrameSizeError)) + } + 0x4 -> + // SETTINGS ACKs never carry params. + case end_stream_or_ack, payload { + True, <<>> -> Ok(Settings(stream_id, True, [])) + True, _payload -> Error(Violation(FrameSizeError)) + False, _payload -> { + use params <- result.try(decode_settings(payload, [])) + Ok(Settings(stream_id:, ack: False, params:)) + } + } + 0x5 -> { + use content <- result.try(strip_padding(padded, payload)) + Ok(PushPromise(stream_id:, payload: content)) + } + 0x6 -> + case payload { + <> -> + Ok(Ping(stream_id:, ack: end_stream_or_ack, opaque_data:)) + _payload -> Error(Violation(FrameSizeError)) + } + 0x7 -> + case stream_id == 0, payload { + False, _payload -> Error(Violation(ProtocolError)) + True, + <<_reserved_bit:1, last_stream_id:31, error_code:32, debug_data:bits>> + -> + Ok(Goaway( + stream_id:, + last_stream_id:, + error_code: decode_error_code(error_code), + debug_data:, + )) + True, _payload -> Error(Violation(FrameSizeError)) + } + 0x8 -> + case payload { + <<_reserved_bit:1, increment:31>> -> + Ok(WindowUpdate(stream_id:, increment:)) + _payload -> Error(Violation(FrameSizeError)) + } + 0x9 -> Ok(Continuation(stream_id:, end_headers:, payload:)) + other -> Ok(Unknown(stream_id, other, payload)) + } +} + +fn strip_padding( + padded: Bool, + payload: BitArray, +) -> Result(BitArray, FrameError) { + case padded { + False -> Ok(payload) + True -> + case payload { + <> -> { + let content_length = bit_array.byte_size(remaining) - pad_length + // A client can claim more padding than bytes actually follow + case content_length >= 0 { + True -> + case remaining { + <> -> + Ok(content) + _remaining -> Error(Violation(FrameSizeError)) + } + False -> Error(Violation(ProtocolError)) + } + } + _payload -> Error(Violation(FrameSizeError)) + } + } +} + +fn strip_priority( + priority: Bool, + stream_id: Int, + payload: BitArray, +) -> Result(BitArray, FrameError) { + case priority { + False -> Ok(payload) + True -> + case payload { + <<_exclusive:1, stream_dependency:31, _weight:8, remaining:bits>> -> + case stream_dependency == stream_id { + True -> Error(Violation(ProtocolError)) + False -> Ok(remaining) + } + _payload -> Error(Violation(FrameSizeError)) + } + } +} + +fn decode_settings( + payload: BitArray, + acc: List(Setting), +) -> Result(List(Setting), FrameError) { + case payload { + <<>> -> Ok(list.reverse(acc)) + <> -> { + use setting <- result.try(decode_setting(id, value)) + decode_settings(remaining, [setting, ..acc]) + } + _payload -> Error(Violation(FrameSizeError)) + } +} + +const max_window_size = 2_147_483_647 + +const min_max_frame_size = 16_384 + +const max_max_frame_size = 16_777_215 + +fn decode_setting(id: Int, value: Int) -> Result(Setting, FrameError) { + case id { + 0x1 -> Ok(HeaderTableSize(value)) + 0x2 -> + case value { + 0 -> Ok(EnablePush(False)) + 1 -> Ok(EnablePush(True)) + _value -> Error(Violation(ProtocolError)) + } + 0x3 -> Ok(MaxConcurrentStreams(value)) + 0x4 -> + case value > max_window_size { + True -> Error(Violation(FlowControlError)) + False -> Ok(InitialWindowSize(value)) + } + 0x5 -> + case value < min_max_frame_size || value > max_max_frame_size { + True -> Error(Violation(ProtocolError)) + False -> Ok(MaxFrameSize(value)) + } + 0x6 -> Ok(MaxHeaderListSize(value)) + other -> Ok(UnknownSetting(other, value)) + } +} + +fn decode_error_code(code: Int) -> ErrorCode { + case code { + 0x0 -> NoError + 0x1 -> ProtocolError + 0x2 -> InternalError + 0x3 -> FlowControlError + 0x4 -> SettingsTimeout + 0x5 -> StreamClosed + 0x6 -> FrameSizeError + 0x7 -> RefusedStream + 0x8 -> Cancel + 0x9 -> CompressionError + 0xa -> ConnectError + 0xb -> EnhanceYourCalm + 0xc -> InadequateSecurity + 0xd -> Http11Required + other -> UnknownErrorCode(other) + } +} + +pub fn encode(frame: Frame) -> BitArray { + let #(type_, flags, payload) = encode_payload(frame) + let length = bit_array.byte_size(payload) + << + length:size(24), + type_:8, + flags:bits, + 0:1, + frame.stream_id:31, + payload:bits, + >> +} + +fn encode_payload(frame: Frame) -> #(Int, BitArray, BitArray) { + case frame { + Data(end_stream:, payload:, ..) -> #( + 0x0, + encode_flags(end_stream, False), + payload, + ) + Headers(end_stream:, end_headers:, payload:, ..) -> #( + 0x1, + encode_flags(end_stream, end_headers), + payload, + ) + Priority(..) -> #(0x2, encode_flags(False, False), <<0:1, 0:31, 0:8>>) + RstStream(error_code:, ..) -> #(0x3, encode_flags(False, False), << + encode_error_code(error_code):32, + >>) + Settings(ack:, params:, ..) -> #( + 0x4, + encode_flags(ack, False), + encode_settings(params), + ) + PushPromise(payload:, ..) -> #(0x5, encode_flags(False, False), payload) + Ping(ack:, opaque_data:, ..) -> #( + 0x6, + encode_flags(ack, False), + opaque_data, + ) + Goaway(last_stream_id:, error_code:, debug_data:, ..) -> #( + 0x7, + encode_flags(False, False), + << + 0:1, + last_stream_id:31, + encode_error_code(error_code):32, + debug_data:bits, + >>, + ) + WindowUpdate(increment:, ..) -> #(0x8, encode_flags(False, False), << + 0:1, + increment:31, + >>) + Continuation(end_headers:, payload:, ..) -> #( + 0x9, + encode_flags(False, end_headers), + payload, + ) + Unknown(frame_type:, payload:, ..) -> #(frame_type, <<0:8>>, payload) + } +} + +fn encode_flags(end_stream_or_ack: Bool, end_headers: Bool) -> BitArray { + <<0:5, bit(end_headers):1, 0:1, bit(end_stream_or_ack):1>> +} + +/// Frames a DATA payload without copying it into the header, so a body already +/// held as a `BytesTree` can be written straight after this. +pub fn encode_data_header( + stream_id: Int, + end_stream: Bool, + length: Int, +) -> BitArray { + << + length:size(24), + 0x0:8, + encode_flags(end_stream, False):bits, + 0:1, + stream_id:31, + >> +} + +fn bit(value: Bool) -> Int { + case value { + True -> 1 + False -> 0 + } +} + +fn encode_settings(params: List(Setting)) -> BitArray { + case params { + [] -> <<>> + [setting, ..rest] -> { + let #(id, value) = encode_setting(setting) + <> + } + } +} + +fn encode_setting(setting: Setting) -> #(Int, Int) { + case setting { + HeaderTableSize(value) -> #(0x1, value) + EnablePush(value) -> #(0x2, bit(value)) + MaxConcurrentStreams(value) -> #(0x3, value) + InitialWindowSize(value) -> #(0x4, value) + MaxFrameSize(value) -> #(0x5, value) + MaxHeaderListSize(value) -> #(0x6, value) + UnknownSetting(id, value) -> #(id, value) + } +} + +fn encode_error_code(code: ErrorCode) -> Int { + case code { + NoError -> 0x0 + ProtocolError -> 0x1 + InternalError -> 0x2 + FlowControlError -> 0x3 + SettingsTimeout -> 0x4 + StreamClosed -> 0x5 + FrameSizeError -> 0x6 + RefusedStream -> 0x7 + Cancel -> 0x8 + CompressionError -> 0x9 + ConnectError -> 0xa + EnhanceYourCalm -> 0xb + InadequateSecurity -> 0xc + Http11Required -> 0xd + UnknownErrorCode(other) -> other + } +} diff --git a/src/ewe/internal/http2/sse.gleam b/src/ewe/internal/http2/sse.gleam new file mode 100644 index 0000000..009381e --- /dev/null +++ b/src/ewe/internal/http2/sse.gleam @@ -0,0 +1,102 @@ +import ewe/internal/connection +import ewe/internal/http2/connection as http2 +import ewe/internal/http2/stream +import ewe/internal/rescue +import ewe/internal/sse +import gleam/bytes_tree +import gleam/erlang/process +import gleam/erlang/reference +import gleam/result +import logging + +/// Runs a Server-Sent Events stream. The response head is already written by +/// the time this is reached so there is nothing to negotiate. +pub fn run( + conn: http2.SseConnection(connection.Body), + on_init: fn(process.Subject(user_message)) -> user_state, + step: fn(connection.SseConnection, user_state, user_message) -> + sse.Step(user_state), + on_close: fn(connection.SseConnection, user_state) -> Nil, +) -> connection.Outcome { + let handle = connection.Http2Sse(conn) + let tag = reference.new() + let subject = process.unsafely_create_subject(process.self(), http2.tag(tag)) + + loop(conn, handle, tag, on_init(subject), step, on_close) +} + +fn loop( + conn: http2.SseConnection(connection.Body), + handle: connection.SseConnection, + tag: reference.Reference, + state: user_state, + step: fn(connection.SseConnection, user_state, user_message) -> + sse.Step(user_state), + on_close: fn(connection.SseConnection, user_state) -> Nil, +) -> connection.Outcome { + case http2.receive_reply(tag) { + // The stream was reset or the connection went. + Error(_interrupted) -> ended(handle, state, on_close, connection.Stopped) + Ok(message) -> + case rescue.handler(fn() { step(handle, state, message) }) { + Error(details) -> crashed(handle, state, on_close, details) + Ok(sse.Proceed(state)) -> loop(conn, handle, tag, state, step, on_close) + Ok(sse.Halt(outcome)) -> { + let outcome = ended(handle, state, on_close, outcome) + + // A stream the handler ended gets a clean close. An abnormal one + // resets. + case outcome { + connection.Stopped -> { + let _sent = stream.finish_response(conn.writer) + Nil + } + connection.StoppedAbnormal(_reason) -> Nil + } + + outcome + } + } + } +} + +/// Every ending runs `on_close` and reports the outcome. +fn ended( + handle: connection.SseConnection, + state: user_state, + on_close: fn(connection.SseConnection, user_state) -> Nil, + outcome: connection.Outcome, +) -> connection.Outcome { + // A bug in `on_close` is still a bug. It just must not kill the process on + // the way out of a stream that already ended. + // TODO: logging here? + let _crashed = rescue.handler(fn() { on_close(handle, state) }) + outcome +} + +/// A crashed handler cannot say what to do next. End the stream for it and +/// reset. +fn crashed( + handle: connection.SseConnection, + state: user_state, + on_close: fn(connection.SseConnection, user_state) -> Nil, + details: String, +) -> connection.Outcome { + logging.log( + logging.Error, + "Caught a crash in the server-sent events handler: " <> details, + ) + + connection.StoppedAbnormal("the handler crashed") + |> ended(handle, state, on_close, _) +} + +pub fn send( + conn: http2.SseConnection(connection.Body), + event: sse.Event, +) -> Result(Nil, http2.Interrupted) { + sse.encode(event) + |> bytes_tree.to_bit_array + |> stream.send_chunk(conn.writer, _) + |> result.replace(Nil) +} diff --git a/src/ewe/internal/http2/stream.gleam b/src/ewe/internal/http2/stream.gleam new file mode 100644 index 0000000..6da0a23 --- /dev/null +++ b/src/ewe/internal/http2/stream.gleam @@ -0,0 +1,162 @@ +import ewe/internal/connection +import ewe/internal/http2/connection as http2 +import ewe/internal/rescue +import gleam/erlang/process +import gleam/erlang/reference +import gleam/http/request +import gleam/http/response +import gleam/result +import logging + +pub fn start( + reply_to: process.Subject(http2.Reply(connection.Body)), + stream_id: Int, + request: request.Request(connection.Connection), + handler: fn(request.Request(connection.Connection)) -> + response.Response(connection.Body), +) -> process.Pid { + use <- process.spawn + + process.trap_exits(True) + + case rescue.handler(fn() { handler(request) }) { + Ok(response) -> deliver(reply_to, stream_id, response) + Error(details) -> crashed(reply_to, stream_id, details) + } +} + +fn crashed( + reply_to: process.Subject(http2.Reply(connection.Body)), + stream_id: Int, + details: String, +) -> Nil { + logging.log(logging.Error, "Caught a crash in the handler: " <> details) + + internal_error(reply_to, stream_id) +} + +fn internal_error( + reply_to: process.Subject(http2.Reply(connection.Body)), + stream_id: Int, +) -> Nil { + response.new(500) + |> response.set_body(connection.Empty) + |> http2.Respond(stream_id, _) + |> process.send(reply_to, _) +} + +fn deliver( + reply_to: process.Subject(http2.Reply(connection.Body)), + stream_id: Int, + response: response.Response(connection.Body), +) -> Nil { + case response.body { + connection.Streaming(connection.StreamingMetadata(handler:)) -> + case begin(reply_to, stream_id, response, http2.Nothing) { + // Nothing was written and nothing can be. No stream left to produce + // a body for. + Error(_interrupted) -> Nil + Ok(writer) -> handler(connection.Http2Writer(writer)) + } + connection.Sse(connection.SseMetadata(handler:)) -> + case begin(reply_to, stream_id, response, http2.SseHeaders) { + Error(_interrupted) -> Nil + Ok(writer) -> + case handler(connection.Http2Sse(http2.SseConnection(writer))) { + connection.Stopped -> Nil + connection.StoppedAbnormal(reason) -> abort(reason) + } + } + // A WebSocket body never gets here. The handshake needs extended CONNECT + // and is refused long before a handler returns one. + connection.Websocket(_metadata) -> { + logging.log( + logging.Error, + "Discarded a WebSocket response: HTTP/2 connections do not carry them", + ) + + internal_error(reply_to, stream_id) + } + connection.Bytes(_tree) + | connection.Text(_text) + | connection.Empty + | connection.File(_file) -> + process.send(reply_to, http2.Respond(stream_id, response)) + } +} + +/// Writes the response head. A handler only gets a writer once the client has +/// something to hang the body off. +fn begin( + reply_to: process.Subject(http2.Reply(connection.Body)), + stream_id: Int, + response: response.Response(connection.Body), + reserved: http2.Reserved, +) -> Result(http2.ResponseWriter(connection.Body), http2.Interrupted) { + let ack_ref = reference.new() + let ack = process.unsafely_create_subject(process.self(), http2.tag(ack_ref)) + + process.send( + reply_to, + http2.WriteHeaders( + stream_id:, + ack:, + status: response.status, + headers: response.headers, + reserved:, + ), + ) + + use _written <- result.map(http2.receive_reply(ack_ref)) + http2.ResponseWriter(connection: reply_to, stream_id:, ack:, ack_ref:) +} + +/// Sends one body chunk. Returns once the connection has it on the wire. An +/// empty chunk gets no frame. +pub fn send_chunk( + writer: http2.ResponseWriter(connection.Body), + chunk: BitArray, +) -> Result(http2.ResponseWriter(connection.Body), http2.Interrupted) { + case chunk { + <<>> -> Ok(writer) + _chunk -> write(writer, chunk, False) + } +} + +pub fn finish_chunk( + writer: http2.ResponseWriter(connection.Body), + chunk: BitArray, +) -> Result(Nil, http2.Interrupted) { + write(writer, chunk, True) |> result.replace(Nil) +} + +pub fn finish_response( + writer: http2.ResponseWriter(connection.Body), +) -> Result(Nil, http2.Interrupted) { + finish_chunk(writer, <<>>) +} + +/// Waiting on the ack is how flow control reaches the handler. The connection +/// answers once the bytes are gone, so a handler outrunning the client blocks +/// here. +fn write( + writer: http2.ResponseWriter(connection.Body), + chunk: BitArray, + end_stream: Bool, +) -> Result(http2.ResponseWriter(connection.Body), http2.Interrupted) { + process.send( + writer.connection, + http2.WriteData( + stream_id: writer.stream_id, + ack: writer.ack, + chunk:, + end_stream:, + ), + ) + + use _written <- result.map(http2.receive_reply(writer.ack_ref)) + writer +} + +@external(erlang, "http2_ffi", "exit_self") +fn abort(reason: String) -> a diff --git a/src/ewe/internal/http2_ffi.erl b/src/ewe/internal/http2_ffi.erl new file mode 100644 index 0000000..e494523 --- /dev/null +++ b/src/ewe/internal/http2_ffi.erl @@ -0,0 +1,103 @@ +-module(http2_ffi). + +-on_load(init/0). +-export([ + init/0, + validate_header_name/2, + validate_header_value/2, + monotonic_ms/0, + split_once/2, + has_forbidden_header_bytes/2, + name_pattern/0, + forbidden_header_pattern/0, + query_pattern/0, + colon_pattern/0, + exit_self/1, + recv_or_exit/1, + recv_or_exit/2 +]). + + +%% Compiles and caches match patterns once at module load. +init() -> + persistent_term:put( + {?MODULE, name}, + binary:compile_pattern([<> || C <- lists:seq($A, $Z)] ++ [<<0>>, <<"\r">>, <<"\n">>]) + ), + persistent_term:put( + {?MODULE, forbidden_header}, + binary:compile_pattern([<<0>>, <<"\r">>, <<"\n">>]) + ), + persistent_term:put({?MODULE, query}, binary:compile_pattern(<<"?">>)), + persistent_term:put({?MODULE, colon}, binary:compile_pattern(<<":">>)), + ok. + +name_pattern() -> persistent_term:get({?MODULE, name}). +forbidden_header_pattern() -> persistent_term:get({?MODULE, forbidden_header}). +query_pattern() -> persistent_term:get({?MODULE, query}). +colon_pattern() -> persistent_term:get({?MODULE, colon}). + +monotonic_ms() -> + erlang:monotonic_time(millisecond). + +exit_self(Reason) -> + erlang:exit(Reason). + +recv_or_exit(Ref) -> + receive + {Ref, Message} -> {ok, Message}; + {'EXIT', _Pid, Reason} -> {error, classify_exit(Reason)} + end. + +recv_or_exit(Ref, Timeout) -> + receive + {Ref, Message} -> {ok, Message}; + {'EXIT', _Pid, Reason} -> {error, classify_exit(Reason)} + after Timeout -> + {error, timed_out} + end. + +classify_exit(<<"stream_reset">>) -> stream_reset; +classify_exit(_Reason) -> connection_closed. + +validate_header_name(_Pattern, <<>>) -> + {error, invalid_utf8}; +validate_header_name(Pattern, Bin) when is_binary(Bin) -> + case ewe_ffi:is_valid_utf8(Bin) of + false -> {error, invalid_utf8}; + true -> classify_name_match(binary:match(Bin, Pattern), Bin) + end; +validate_header_name(_Pattern, _Bits) -> + {error, invalid_utf8}. + +%% We do one scan for both `has an uppercase letter` and `has a forbidden byte`. +%% We only need to know which kind of bad byte it was once we've found one. +classify_name_match(nomatch, Bin) -> + {ok, Bin}; +classify_name_match({Pos, _Len}, Bin) -> + case binary:at(Bin, Pos) of + C when C >= $A, C =< $Z -> {error, uppercase_header_name}; + _Byte -> {error, malformed_header_bytes} + end. + +validate_header_value(Pattern, Bin) when is_binary(Bin) -> + case ewe_ffi:is_valid_utf8(Bin) of + false -> + {error, invalid_utf8}; + true -> + case has_forbidden_header_bytes(Pattern, Bin) of + true -> {error, malformed_header_bytes}; + false -> {ok, Bin} + end + end; +validate_header_value(_Pattern, _Bits) -> + {error, invalid_utf8}. + +has_forbidden_header_bytes(Pattern, Bin) -> + binary:match(Bin, Pattern) =/= nomatch. + +split_once(Bin, Pattern) -> + case binary:split(Bin, Pattern) of + [Before, After] -> {ok, {Before, After}}; + _Parts -> {error, nil} + end. diff --git a/test/ewe/internal/http2/body_test.gleam b/test/ewe/internal/http2/body_test.gleam new file mode 100644 index 0000000..8012113 --- /dev/null +++ b/test/ewe/internal/http2/body_test.gleam @@ -0,0 +1,183 @@ +import ewe/internal/connection as ewe_connection +import ewe/internal/http2/body +import ewe/internal/http2/connection as http2 +import gleam/erlang/process +import gleam/option + +fn fill_mailbox( + conn_subject: process.Subject(http2.Reply(ewe_connection.Body)), + script: List(http2.BodyEvent), +) -> Nil { + case script { + [] -> Nil + [event, ..remaining] -> { + let assert http2.ReadBody(_stream_id, reply_to) = + process.receive_forever(conn_subject) + process.send(reply_to, event) + fill_mailbox(conn_subject, remaining) + } + } +} + +fn mock_connection( + script: List(http2.BodyEvent), +) -> http2.Connection(ewe_connection.Body) { + let init_subject = process.new_subject() + process.spawn(fn() { + let conn_subject = process.new_subject() + process.send(init_subject, conn_subject) + fill_mailbox(conn_subject, script) + }) + let conn_subject = process.receive_forever(init_subject) + http2.Connection( + connection: conn_subject, + stream_id: 1, + has_body: True, + pending: <<>>, + pending_trailers: option.None, + read: 0, + body_read_timeout: 1000, + peer: Error(Nil), + ) +} + +pub fn read_body_accumulates_chunks_until_done_test() { + let conn = + mock_connection([ + http2.ChunkEvent(<<"abc":utf8>>), + http2.ChunkEvent(<<"def":utf8>>), + http2.DoneEvent([]), + ]) + + assert body.read_body(conn, 100) == Ok(#(<<"abcdef":utf8>>, [])) +} + +pub fn read_body_returns_trailers_test() { + let conn = + mock_connection([ + http2.ChunkEvent(<<"abc":utf8>>), + http2.DoneEvent([#("x-checksum", "deadbeef")]), + ]) + + assert body.read_body(conn, 100) + == Ok(#(<<"abc":utf8>>, [#("x-checksum", "deadbeef")])) +} + +pub fn read_body_collapses_last_chunk_into_one_round_trip_test() { + let conn = mock_connection([http2.LastChunkEvent(<<"abcdef":utf8>>, [])]) + + assert body.read_body(conn, 100) == Ok(#(<<"abcdef":utf8>>, [])) +} + +pub fn read_body_collapsed_last_chunk_carries_trailers_test() { + let conn = + mock_connection([ + http2.LastChunkEvent(<<"abc":utf8>>, [#("x-checksum", "deadbeef")]), + ]) + + assert body.read_body(conn, 100) + == Ok(#(<<"abc":utf8>>, [#("x-checksum", "deadbeef")])) +} + +pub fn read_body_chunk_collapses_final_lump_into_one_round_trip_test() { + let conn = mock_connection([http2.LastChunkEvent(<<"abcdefghij":utf8>>, [])]) + + let assert Ok(body.Chunk(first, conn)) = + body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + assert first == <<"abcd":utf8>> + + let assert Ok(body.Chunk(second, conn)) = + body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + assert second == <<"efgh":utf8>> + + let assert Ok(body.Chunk(third, conn)) = + body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + assert third == <<"ij":utf8>> + + assert body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + == Ok(body.Done([])) +} + +pub fn read_body_empty_body_is_done_immediately_test() { + let conn = mock_connection([http2.DoneEvent([])]) + + assert body.read_body(conn, 100) == Ok(#(<<>>, [])) +} + +pub fn read_body_no_body_skips_round_trip_test() { + let conn = + http2.Connection( + connection: process.new_subject(), + stream_id: 1, + has_body: False, + pending: <<>>, + pending_trailers: option.None, + read: 0, + body_read_timeout: 1000, + peer: Error(Nil), + ) + + assert body.read_body(conn, 100) == Ok(#(<<>>, [])) +} + +pub fn read_body_stops_when_over_limit_test() { + let chunk = <<0:size({ 30 * 8 })>> + let conn = mock_connection([http2.ChunkEvent(chunk), http2.ChunkEvent(chunk)]) + + assert body.read_body(conn, 40) == Error(body.BodyTooLarge) +} + +pub fn read_body_chunk_splits_pulled_lump_without_extra_message_test() { + let conn = + mock_connection([ + http2.ChunkEvent(<<"abcdefghij":utf8>>), + http2.DoneEvent([]), + ]) + + let assert Ok(body.Chunk(first, conn)) = + body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + assert first == <<"abcd":utf8>> + + let assert Ok(body.Chunk(second, conn)) = + body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + assert second == <<"efgh":utf8>> + + let assert Ok(body.Chunk(third, conn)) = + body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + assert third == <<"ij":utf8>> + + assert body.read_body_chunk(conn, max_chunk_bytes: 4, limit: 100) + == Ok(body.Done([])) +} + +pub fn read_body_chunk_returns_whole_lump_when_under_cap_test() { + let conn = + mock_connection([ + http2.ChunkEvent(<<"abc":utf8>>), + http2.DoneEvent([#("x-checksum", "deadbeef")]), + ]) + + let assert Ok(body.Chunk(data, conn)) = + body.read_body_chunk(conn, max_chunk_bytes: 100, limit: 100) + assert data == <<"abc":utf8>> + + assert body.read_body_chunk(conn, max_chunk_bytes: 100, limit: 100) + == Ok(body.Done([#("x-checksum", "deadbeef")])) +} + +pub fn read_body_chunk_no_body_skips_round_trip_test() { + let conn = + http2.Connection( + connection: process.new_subject(), + stream_id: 1, + has_body: False, + pending: <<>>, + pending_trailers: option.None, + read: 0, + body_read_timeout: 1000, + peer: Error(Nil), + ) + + assert body.read_body_chunk(conn, max_chunk_bytes: 100, limit: 100) + == Ok(body.Done([])) +} diff --git a/test/ewe/internal/http2/connection_test.gleam b/test/ewe/internal/http2/connection_test.gleam new file mode 100644 index 0000000..a396aaa --- /dev/null +++ b/test/ewe/internal/http2/connection_test.gleam @@ -0,0 +1,1103 @@ +import alpacki +import ewe/internal/http2 as connection +import ewe/internal/http2/connection as http2 +import ewe/internal/http2/frame +import gleam/bit_array +import gleam/bytes_tree +import gleam/dict +import gleam/erlang/process +import gleam/http +import gleam/int +import gleam/list +import gleam/option.{None, Some} + +fn pending_stream(send_window: Int, pending: BitArray) -> connection.Stream { + connection.Stream( + status: connection.Flushing, + send_window:, + pending: connection.PendingBytes(pending), + pending_end_stream: True, + write_ack: None, + recv_window: 65_535, + recv_buffer: bytes_tree.new(), + request_half_closed: False, + parked_reader: None, + content_length: None, + body_bytes_received: 0, + trailers: [], + ) +} + +fn inbound_stream( + recv_window: Int, + parked_reader: option.Option(process.Subject(http2.BodyEvent)), +) -> connection.Stream { + connection.Stream( + status: connection.Flushing, + send_window: 65_535, + pending: connection.PendingBytes(<<>>), + pending_end_stream: True, + write_ack: None, + recv_window:, + recv_buffer: bytes_tree.new(), + request_half_closed: False, + parked_reader:, + content_length: None, + body_bytes_received: 0, + trailers: [], + ) +} + +fn encode(fields: List(alpacki.HeaderField)) -> BitArray { + let #(payload, _) = + alpacki.encode_header_block(fields, alpacki.new_dynamic(4096), False) + payload +} + +fn method_get() -> alpacki.HeaderField { + alpacki.HeaderField( + <<":method":utf8>>, + <<"GET":utf8>>, + alpacki.WithoutIndexing, + ) +} + +fn minimal_headers() -> List(alpacki.HeaderField) { + [ + method_get(), + alpacki.HeaderField( + <<":scheme":utf8>>, + <<"https":utf8>>, + alpacki.WithoutIndexing, + ), + alpacki.HeaderField( + <<":authority":utf8>>, + <<"example.com":utf8>>, + alpacki.WithoutIndexing, + ), + alpacki.HeaderField(<<":path":utf8>>, <<"/":utf8>>, alpacki.WithoutIndexing), + ] +} + +fn custom_field() -> alpacki.HeaderField { + alpacki.HeaderField( + <<"x-custom":utf8>>, + <<"some-longer-header-value-here":utf8>>, + alpacki.WithoutIndexing, + ) +} + +fn headers_with_content_length(length: String) -> List(alpacki.HeaderField) { + list.append(minimal_headers(), [ + alpacki.HeaderField( + <<"content-length":utf8>>, + <>, + alpacki.WithoutIndexing, + ), + ]) +} + +pub fn headers_single_frame_completes_test() { + let payload = encode(minimal_headers()) + let result = + connection.handle_headers(connection.test_state(), 1, True, True, payload) + let assert connection.Proceed(state) = result + assert state.header_assembly == None + assert state.highest_client_stream_id_seen == 1 +} + +pub fn headers_end_stream_marks_new_stream_half_closed_test() { + let payload = encode(minimal_headers()) + let result = + connection.handle_headers(connection.test_state(), 1, True, True, payload) + let assert connection.Proceed(state) = result + let assert Ok(entry) = dict.get(state.streams, 1) + assert entry.request_half_closed == True +} + +pub fn headers_without_end_stream_leaves_stream_open_test() { + let payload = encode(minimal_headers()) + let result = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + let assert connection.Proceed(state) = result + let assert Ok(entry) = dict.get(state.streams, 1) + assert entry.request_half_closed == False +} + +pub fn headers_even_stream_id_is_protocol_error_test() { + let payload = encode(minimal_headers()) + let result = + connection.handle_headers(connection.test_state(), 2, True, True, payload) + assert result == connection.Terminate(Some(frame.ProtocolError)) +} + +pub fn headers_reused_stream_id_on_half_closed_remote_is_stream_error_test() { + let payload = encode(minimal_headers()) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 3, True, True, payload) + + let result = connection.handle_headers(state, 3, True, True, payload) + let assert connection.RejectStream(_, 3, frame.StreamClosed) = result +} + +pub fn headers_decreasing_stream_id_is_protocol_error_test() { + let payload = encode(minimal_headers()) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 5, True, True, payload) + + let result = connection.handle_headers(state, 3, True, True, payload) + assert result == connection.Terminate(Some(frame.ProtocolError)) +} + +pub fn headers_increasing_stream_id_after_reject_still_advances_test() { + let malformed = encode([method_get()]) + let assert connection.RejectStream(state, 1, _) = + connection.handle_headers(connection.test_state(), 1, True, True, malformed) + + let payload = encode(minimal_headers()) + let result = connection.handle_headers(state, 1, True, True, payload) + assert result == connection.Terminate(Some(frame.ProtocolError)) +} + +pub fn headers_while_draining_refuses_new_stream_test() { + let state = connection.State(..connection.test_state(), draining: True) + let payload = encode(minimal_headers()) + + let result = connection.handle_headers(state, 1, True, True, payload) + + let assert connection.RejectStream(state, 1, frame.RefusedStream) = result + assert dict.get(state.streams, 1) == Error(Nil) +} + +pub fn headers_without_end_headers_starts_assembly_test() { + let payload = encode([method_get()]) + let result = + connection.handle_headers(connection.test_state(), 1, True, False, payload) + let assert connection.Proceed(state) = result + assert state.header_assembly + == Some(connection.HeaderAssembly(1, True, 1, payload, False)) +} + +pub fn oversized_headers_frame_is_enhance_your_calm_test() { + let huge = <<0:size({ 65_537 * 8 })>> + let result = + connection.handle_headers(connection.test_state(), 1, True, True, huge) + assert result == connection.Terminate(Some(frame.EnhanceYourCalm)) +} + +pub fn continuation_completes_split_block_test() { + let payload = encode(list.append(minimal_headers(), [custom_field()])) + let assert <> = payload + + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, True, False, first) + let assert Some(assembly) = state.header_assembly + + let result = + connection.handle_continuation( + state, + assembly, + frame.Continuation(1, True, second), + ) + let assert connection.Proceed(final_state) = result + assert final_state.header_assembly == None +} + +pub fn continuation_wrong_stream_is_protocol_error_test() { + let assembly = connection.HeaderAssembly(1, True, 1, <<>>, False) + let result = + connection.handle_continuation( + connection.test_state(), + assembly, + frame.Continuation(2, True, <<>>), + ) + assert result == connection.Terminate(Some(frame.ProtocolError)) +} + +pub fn non_continuation_mid_assembly_is_protocol_error_test() { + let assembly = connection.HeaderAssembly(1, True, 1, <<>>, False) + let result = + connection.handle_continuation( + connection.test_state(), + assembly, + frame.Ping(0, False, <<0, 0, 0, 0, 0, 0, 0, 0>>), + ) + assert result == connection.Terminate(Some(frame.ProtocolError)) +} + +pub fn add_fragment_within_limits_test() { + let assembly = connection.HeaderAssembly(1, True, 1, <<"a":utf8>>, False) + let assert Ok(updated) = + connection.add_fragment(assembly, <<"b":utf8>>, http2.default_config()) + assert updated == connection.HeaderAssembly(1, True, 2, <<"ab":utf8>>, False) +} + +pub fn add_fragment_over_count_cap_is_enhance_your_calm_test() { + let assembly = connection.HeaderAssembly(1, True, 100, <<>>, False) + let result = connection.add_fragment(assembly, <<>>, http2.default_config()) + assert result == Error(frame.EnhanceYourCalm) +} + +pub fn add_fragment_over_byte_cap_is_enhance_your_calm_test() { + let assembly = + connection.HeaderAssembly(1, True, 1, <<0:size({ 65_536 * 8 })>>, False) + let result = + connection.add_fragment(assembly, <<"x":utf8>>, http2.default_config()) + assert result == Error(frame.EnhanceYourCalm) +} + +pub fn complete_header_block_invalid_hpack_is_compression_error_test() { + let assembly = connection.HeaderAssembly(1, True, 1, <<0xff, 0xff>>, False) + let result = + connection.complete_header_block(connection.test_state(), assembly) + assert result == connection.Terminate(Some(frame.CompressionError)) +} + +pub fn complete_header_block_oversized_list_is_enhance_your_calm_test() { + let big_value = <<0:size({ 20_000 * 8 })>> + let field = + alpacki.HeaderField(<<"x":utf8>>, big_value, alpacki.WithoutIndexing) + let assembly = connection.HeaderAssembly(1, True, 1, encode([field]), False) + let config = + http2.Config(..http2.default_config(), max_header_list_size: Some(16_384)) + let state = connection.State(..connection.test_state(), config:) + let result = connection.complete_header_block(state, assembly) + assert result == connection.Terminate(Some(frame.EnhanceYourCalm)) +} + +pub fn complete_header_block_oversized_list_unlimited_by_default_test() { + let big_value = <<0:size({ 20_000 * 8 })>> + let field = + alpacki.HeaderField(<<"x":utf8>>, big_value, alpacki.WithoutIndexing) + let assembly = connection.HeaderAssembly(1, True, 1, encode([field]), False) + let result = + connection.complete_header_block(connection.test_state(), assembly) + assert result != connection.Terminate(Some(frame.EnhanceYourCalm)) +} + +pub fn build_request_minimal_valid_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + ] + let assert Ok(#(request, None)) = + connection.build_request(headers, Nil, connection.header_patterns()) + assert request.method == http.Get + assert request.scheme == http.Https + assert request.host == "example.com" + assert request.port == None + assert request.path == "/" + assert request.query == None + assert request.headers == [] +} + +pub fn build_request_with_query_and_port_test() { + let headers = [ + #(<<":method":utf8>>, <<"POST":utf8>>), + #(<<":scheme":utf8>>, <<"http":utf8>>), + #(<<":authority":utf8>>, <<"example.com:8080":utf8>>), + #(<<":path":utf8>>, <<"/search?q=1":utf8>>), + #(<<"x-custom":utf8>>, <<"value":utf8>>), + ] + let assert Ok(#(request, None)) = + connection.build_request(headers, Nil, connection.header_patterns()) + assert request.method == http.Post + assert request.host == "example.com" + assert request.port == Some(8080) + assert request.path == "/search" + assert request.query == Some("q=1") + assert request.headers == [#("x-custom", "value")] +} + +pub fn build_request_missing_pseudo_header_test() { + let headers = [#(<<":method":utf8>>, <<"GET":utf8>>)] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.MissingPseudoHeader) +} + +pub fn build_request_duplicate_pseudo_header_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":method":utf8>>, <<"POST":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.DuplicatePseudoHeader) +} + +pub fn build_request_pseudo_after_regular_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<"x-custom":utf8>>, <<"value":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.PseudoHeaderAfterRegular) +} + +pub fn build_request_unknown_pseudo_header_test() { + let headers = [ + #(<<":bogus":utf8>>, <<"x":utf8>>), + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.UnknownPseudoHeader) +} + +pub fn build_request_invalid_scheme_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"ftp":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.InvalidScheme) +} + +pub fn build_request_invalid_authority_port_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com:abc":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.InvalidAuthority) +} + +pub fn build_request_empty_path_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.InvalidPath) +} + +pub fn build_request_empty_header_name_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + #(<<>>, <<"value":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.EmptyHeaderName) +} + +pub fn build_request_uppercase_header_name_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + #(<<"X-Custom":utf8>>, <<"value":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.UppercaseHeaderName) +} + +pub fn build_request_connection_specific_header_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + #(<<"connection":utf8>>, <<"keep-alive":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.ConnectionSpecificHeader) +} + +pub fn build_request_te_trailers_allowed_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + #(<<"te":utf8>>, <<"trailers":utf8>>), + ] + let assert Ok(#(request, None)) = + connection.build_request(headers, Nil, connection.header_patterns()) + assert request.headers == [#("te", "trailers")] +} + +pub fn build_request_te_non_trailers_rejected_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + #(<<"te":utf8>>, <<"gzip":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.ConnectionSpecificHeader) +} + +pub fn build_request_invalid_utf8_header_value_test() { + let headers = [ + #(<<":method":utf8>>, <<"GET":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + #(<<"x-custom":utf8>>, <<0xff, 0xfe>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.InvalidUtf8) +} + +pub fn build_request_invalid_method_test() { + let headers = [ + #(<<":method":utf8>>, <<"bad method":utf8>>), + #(<<":scheme":utf8>>, <<"https":utf8>>), + #(<<":authority":utf8>>, <<"example.com":utf8>>), + #(<<":path":utf8>>, <<"/":utf8>>), + ] + assert connection.build_request(headers, Nil, connection.header_patterns()) + == Error(connection.InvalidMethod) +} + +pub fn complete_header_block_updates_dynamic_table_test() { + let field = + alpacki.HeaderField( + <<"x-custom":utf8>>, + <<"value":utf8>>, + alpacki.WithIndexing, + ) + let payload = encode(list.append(minimal_headers(), [field])) + let assembly = connection.HeaderAssembly(1, True, 1, payload, False) + let result = + connection.complete_header_block(connection.test_state(), assembly) + let assert connection.Proceed(state) = result + assert alpacki.dynamic_length(state.hpack_decoder) == 1 +} + +pub fn complete_header_block_invalid_request_rejects_stream_test() { + let assembly = + connection.HeaderAssembly(1, True, 1, encode([method_get()]), False) + let result = + connection.complete_header_block(connection.test_state(), assembly) + let assert connection.RejectStream(_, stream_id, code) = result + assert stream_id == 1 + assert code == frame.ProtocolError +} + +pub fn client_reset_under_threshold_continues_test() { + let state = + connection.State( + ..connection.test_state(), + highest_client_stream_id_seen: 1, + ) + let result = connection.handle_client_reset(state, 1) + let assert connection.Proceed(_) = result +} + +pub fn client_reset_on_idle_stream_is_protocol_error_test() { + let result = connection.handle_client_reset(connection.test_state(), 1) + assert result == connection.Terminate(Some(frame.ProtocolError)) +} + +pub fn rapid_reset_trips_enhance_your_calm_test() { + let state = + int.range( + from: 1, + to: 101, + with: connection.test_state(), + run: fn(state, stream_id) { + let state = + connection.State(..state, highest_client_stream_id_seen: stream_id) + let assert connection.Proceed(state) = + connection.handle_client_reset(state, stream_id) + state + }, + ) + + let result = connection.handle_client_reset(state, 101) + let assert connection.Terminate(Some(frame.EnhanceYourCalm)) = result +} + +pub fn append_header_frames_splits_into_continuation_test() { + let block = <<0:size(120)>> + let out = + connection.append_header_frames(bytes_tree.new(), 1, True, block, 10) + |> bytes_tree.to_bit_array + + let assert Ok(#(frame.Headers(1, True, False, first_chunk), rest)) = + frame.decode(out, 16_384) + let assert Ok(#(frame.Continuation(1, True, second_chunk), rest)) = + frame.decode(rest, 16_384) + + assert bit_array.byte_size(first_chunk) == 10 + assert bit_array.byte_size(second_chunk) == 5 + assert rest == <<>> +} + +pub fn flush_stream_full_drain_test() { + let state = connection.test_state() + let entry = pending_stream(100, <<"hello":utf8>>) + + let assert connection.FlushAccumulated(state, out, wrote) = + connection.do_flush_stream(state, 1, entry, bytes_tree.new(), False) + + assert wrote == True + let assert Ok(#(frame.Data(1, True, <<"hello":utf8>>, 5), rest)) = + frame.decode(bytes_tree.to_bit_array(out), 16_384) + assert rest == <<>> + assert dict.get(state.streams, 1) == Error(Nil) + assert state.conn_send_window == 65_535 - 5 +} + +pub fn flush_stream_partial_drain_blocks_on_stream_window_test() { + let state = connection.test_state() + let entry = pending_stream(3, <<"hello":utf8>>) + + let assert connection.FlushAccumulated(state, out, wrote) = + connection.do_flush_stream(state, 1, entry, bytes_tree.new(), False) + + assert wrote == True + let assert Ok(#(frame.Data(1, False, <<"hel":utf8>>, 3), rest)) = + frame.decode(bytes_tree.to_bit_array(out), 16_384) + assert rest == <<>> + let assert Ok(remaining) = dict.get(state.streams, 1) + assert remaining.status == connection.Flushing + assert remaining.send_window == 0 + assert remaining.pending == connection.PendingBytes(<<"lo":utf8>>) + assert state.conn_send_window == 65_535 - 3 +} + +pub fn flush_stream_zero_window_defers_everything_test() { + let state = connection.State(..connection.test_state(), conn_send_window: 0) + let entry = pending_stream(100, <<"hello":utf8>>) + + let assert connection.FlushAccumulated(state, out, wrote) = + connection.do_flush_stream(state, 1, entry, bytes_tree.new(), False) + + assert wrote == False + assert out == bytes_tree.new() + let assert Ok(unchanged) = dict.get(state.streams, 1) + assert unchanged == entry +} + +pub fn flush_stream_sends_multiple_frames_in_one_call_when_window_allows_test() { + let state = connection.test_state() + let body = <<0:size({ 40_000 * 8 })>> + let entry = pending_stream(100_000, body) + + let assert connection.FlushAccumulated(state, out, wrote) = + connection.do_flush_stream(state, 1, entry, bytes_tree.new(), False) + + assert wrote == True + let bytes = bytes_tree.to_bit_array(out) + + let assert Ok(#(frame.Data(1, False, chunk_1, _), rest)) = + frame.decode(bytes, 16_384) + assert bit_array.byte_size(chunk_1) == 16_384 + + let assert Ok(#(frame.Data(1, False, chunk_2, _), rest)) = + frame.decode(rest, 16_384) + assert bit_array.byte_size(chunk_2) == 16_384 + + let assert Ok(#(frame.Data(1, True, chunk_3, _), rest)) = + frame.decode(rest, 16_384) + assert bit_array.byte_size(chunk_3) == 40_000 - 16_384 - 16_384 + + assert rest == <<>> + assert dict.get(state.streams, 1) == Error(Nil) + assert state.conn_send_window == 65_535 - 40_000 +} + +pub fn adjust_stream_windows_applies_delta_to_all_streams_test() { + let entry_a = pending_stream(100, <<"a":utf8>>) + let entry_b = pending_stream(200, <<"b":utf8>>) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry_a), #(3, entry_b)]), + ) + + let state = connection.adjust_stream_windows(state, -50) + + let assert Ok(a) = dict.get(state.streams, 1) + let assert Ok(b) = dict.get(state.streams, 3) + assert a.send_window == 50 + assert b.send_window == 150 +} + +pub fn client_reset_on_flushing_stream_removes_it_test() { + let entry = pending_stream(10, <<"tail":utf8>>) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + ) + + let result = connection.handle_client_reset(state, 1) + let assert connection.Proceed(state) = result + assert dict.get(state.streams, 1) == Error(Nil) +} + +pub fn client_reset_on_computing_stream_keeps_stream_pids_test() { + let pid = process.spawn_unlinked(fn() { process.sleep_forever() }) + let entry = + connection.Stream( + ..pending_stream(65_535, <<>>), + status: connection.Computing(pid), + ) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + stream_pids: dict.from_list([#(pid, 1)]), + highest_client_stream_id_seen: 1, + ) + + let result = connection.handle_client_reset(state, 1) + let assert connection.Proceed(state) = result + + assert dict.get(state.streams, 1) == Error(Nil) + assert dict.get(state.stream_pids, pid) == Ok(1) +} + +pub fn handle_data_accumulates_into_buffer_test() { + let entry = inbound_stream(65_535, None) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + conn_recv_window: 2_097_152, + ) + + let assert connection.Proceed(state) = + connection.handle_data(state, 1, False, <<"abc":utf8>>, 3) + + let assert Ok(updated) = dict.get(state.streams, 1) + assert bytes_tree.to_bit_array(updated.recv_buffer) == <<"abc":utf8>> + assert updated.recv_window == 65_535 - 3 + assert updated.request_half_closed == False + assert state.conn_recv_window == 2_097_152 - 3 +} + +pub fn handle_data_padded_frame_counts_full_wire_size_test() { + let entry = inbound_stream(65_535, None) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + conn_recv_window: 2_097_152, + ) + + let assert connection.Proceed(state) = + connection.handle_data(state, 1, False, <<"abc":utf8>>, 10) + + let assert Ok(updated) = dict.get(state.streams, 1) + assert bytes_tree.to_bit_array(updated.recv_buffer) == <<"abc":utf8>> + assert updated.recv_window == 65_535 - 10 + assert state.conn_recv_window == 2_097_152 - 10 +} + +pub fn handle_data_end_stream_sets_half_closed_test() { + let entry = inbound_stream(65_535, None) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + conn_recv_window: 2_097_152, + ) + + let assert connection.Proceed(state) = + connection.handle_data(state, 1, True, <<>>, 0) + + let assert Ok(updated) = dict.get(state.streams, 1) + assert updated.request_half_closed == True +} + +pub fn handle_data_delivers_directly_to_parked_reader_test() { + let reply_to = process.new_subject() + let entry = inbound_stream(2_097_152, Some(reply_to)) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + conn_recv_window: 2_097_152, + ) + + let assert connection.Proceed(state) = + connection.handle_data(state, 1, False, <<"hi":utf8>>, 2) + + let assert Ok(http2.ChunkEvent(<<"hi":utf8>>)) = + process.receive(reply_to, 100) + let assert Ok(updated) = dict.get(state.streams, 1) + assert updated.parked_reader == None + assert bytes_tree.to_bit_array(updated.recv_buffer) == <<>> +} + +pub fn handle_data_delivered_to_parked_reader_grants_credit_test() { + let reply_to = process.new_subject() + let entry = inbound_stream(100_000, Some(reply_to)) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + ) + let chunk = <<0:size({ 40_000 * 8 })>> + + let result = connection.handle_data(state, 1, False, chunk, 40_000) + + let assert connection.ProceedWithOutbound(state, out) = result + let assert Ok(http2.ChunkEvent(delivered)) = process.receive(reply_to, 100) + assert delivered == chunk + + let bytes = bytes_tree.to_bit_array(out) + let assert Ok(#(frame.WindowUpdate(1, 2_037_152), rest)) = + frame.decode(bytes, 16_384) + let assert Ok(#(frame.WindowUpdate(0, 2_071_617), rest)) = + frame.decode(rest, 16_384) + assert rest == <<>> + + let assert Ok(updated) = dict.get(state.streams, 1) + assert updated.recv_window == 2_097_152 +} + +pub fn handle_data_delivers_last_chunk_to_parked_reader_test() { + let reply_to = process.new_subject() + let entry = inbound_stream(2_097_152, Some(reply_to)) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + conn_recv_window: 2_097_152, + ) + + let assert connection.Proceed(_state) = + connection.handle_data(state, 1, True, <<"hi":utf8>>, 2) + + let assert Ok(http2.LastChunkEvent(<<"hi":utf8>>, [])) = + process.receive(reply_to, 100) +} + +pub fn handle_data_sends_done_to_parked_reader_on_empty_end_stream_test() { + let reply_to = process.new_subject() + let entry = inbound_stream(2_097_152, Some(reply_to)) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + conn_recv_window: 2_097_152, + ) + + let assert connection.Proceed(_state) = + connection.handle_data(state, 1, True, <<>>, 0) + + let assert Ok(http2.DoneEvent([])) = process.receive(reply_to, 100) +} + +pub fn handle_data_stream_window_violation_rejects_stream_test() { + let entry = inbound_stream(5, None) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + ) + + let result = connection.handle_data(state, 1, False, <<0:size(80)>>, 10) + + let assert connection.RejectStream(state, 1, frame.FlowControlError) = result + assert state.conn_recv_window == 65_535 - 10 +} + +pub fn handle_data_conn_window_violation_terminates_test() { + let entry = inbound_stream(65_535, None) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + conn_recv_window: 5, + highest_client_stream_id_seen: 1, + ) + + let result = connection.handle_data(state, 1, False, <<0:size(80)>>, 10) + + let assert connection.Terminate(Some(frame.FlowControlError)) = result +} + +pub fn handle_data_on_idle_stream_is_protocol_error_test() { + let state = connection.test_state() + + let result = connection.handle_data(state, 3, False, <<"x":utf8>>, 1) + + let assert connection.Terminate(Some(frame.ProtocolError)) = result +} + +pub fn handle_data_on_already_closed_stream_is_connection_error_test() { + let state = + connection.State( + ..connection.test_state(), + highest_client_stream_id_seen: 5, + ) + + let result = connection.handle_data(state, 3, False, <<"x":utf8>>, 1) + + assert result == connection.Terminate(Some(frame.StreamClosed)) +} + +pub fn handle_data_buffered_without_reader_still_credits_conn_window_test() { + let entry = inbound_stream(100_000, None) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + ) + let chunk = <<0:size({ 40_000 * 8 })>> + + let result = connection.handle_data(state, 1, False, chunk, 40_000) + + let assert connection.ProceedWithOutbound(state, out) = result + let bytes = bytes_tree.to_bit_array(out) + let assert Ok(#(frame.WindowUpdate(0, 2_071_617), rest)) = + frame.decode(bytes, 16_384) + assert rest == <<>> + assert state.conn_recv_window == 2_097_152 + + let assert Ok(updated) = dict.get(state.streams, 1) + assert updated.recv_window == 100_000 - 40_000 +} + +pub fn handle_data_on_already_closed_stream_with_large_chunk_is_connection_error_test() { + let state = + connection.State( + ..connection.test_state(), + highest_client_stream_id_seen: 5, + ) + let chunk = <<0:size({ 40_000 * 8 })>> + + let result = connection.handle_data(state, 3, False, chunk, 40_000) + + assert result == connection.Terminate(Some(frame.StreamClosed)) +} + +pub fn handle_data_on_already_closed_stream_can_exceed_conn_window_test() { + let state = + connection.State( + ..connection.test_state(), + highest_client_stream_id_seen: 5, + conn_recv_window: 5, + ) + + let result = connection.handle_data(state, 3, False, <<0:size(80)>>, 10) + + assert result == connection.Terminate(Some(frame.FlowControlError)) +} + +pub fn handle_data_with_stream_id_zero_is_protocol_error_test() { + let state = connection.test_state() + + let result = connection.handle_data(state, 0, False, <<"x":utf8>>, 1) + + assert result == connection.Terminate(Some(frame.ProtocolError)) +} + +pub fn handle_data_after_half_closed_remote_rejects_stream_test() { + let entry = + connection.Stream(..inbound_stream(65_535, None), request_half_closed: True) + let state = + connection.State( + ..connection.test_state(), + streams: dict.from_list([#(1, entry)]), + highest_client_stream_id_seen: 1, + ) + + let result = connection.handle_data(state, 1, False, <<"x":utf8>>, 1) + + let assert connection.RejectStream(state, 1, frame.StreamClosed) = result + assert state.conn_recv_window == 65_535 - 1 +} + +pub fn stream_recv_credit_above_low_water_mark_does_not_emit_test() { + let entry = inbound_stream(300_000, None) + + let #(entry, increment) = + connection.stream_recv_credit(entry, http2.default_config()) + + assert increment == 0 + assert entry.recv_window == 300_000 +} + +pub fn stream_recv_credit_at_low_water_mark_refills_to_high_test() { + let entry = inbound_stream(262_144, None) + + let #(entry, increment) = + connection.stream_recv_credit(entry, http2.default_config()) + + assert increment == 2_097_152 - 262_144 + assert entry.recv_window == 2_097_152 +} + +pub fn stream_recv_credit_below_low_water_mark_refills_to_high_test() { + let entry = inbound_stream(1000, None) + + let #(entry, increment) = + connection.stream_recv_credit(entry, http2.default_config()) + + assert increment == 2_097_152 - 1000 + assert entry.recv_window == 2_097_152 +} + +pub fn conn_recv_credit_above_low_water_mark_does_not_emit_test() { + let state = + connection.State(..connection.test_state(), conn_recv_window: 300_000) + + let #(state, increment) = connection.conn_recv_credit(state) + + assert increment == 0 + assert state.conn_recv_window == 300_000 +} + +pub fn conn_recv_credit_below_low_water_mark_refills_to_high_test() { + let state = + connection.State(..connection.test_state(), conn_recv_window: 1000) + + let #(state, increment) = connection.conn_recv_credit(state) + + assert increment == 2_097_152 - 1000 + assert state.conn_recv_window == 2_097_152 +} + +pub fn trailers_after_data_marks_half_closed_test() { + let payload = encode(minimal_headers()) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let trailer_payload = encode([custom_field()]) + let result = connection.handle_headers(state, 1, True, True, trailer_payload) + + let assert connection.Proceed(state) = result + let assert Ok(entry) = dict.get(state.streams, 1) + assert entry.request_half_closed == True +} + +pub fn trailers_with_pseudo_header_is_protocol_error_test() { + let payload = encode(minimal_headers()) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let trailer_payload = encode([method_get()]) + let result = connection.handle_headers(state, 1, True, True, trailer_payload) + + let assert connection.RejectStream(_, 1, frame.ProtocolError) = result +} + +pub fn trailers_without_end_stream_is_protocol_error_test() { + let payload = encode(minimal_headers()) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let trailer_payload = encode([custom_field()]) + let result = connection.handle_headers(state, 1, False, True, trailer_payload) + + let assert connection.RejectStream(_, 1, frame.ProtocolError) = result +} + +pub fn trailers_are_delivered_to_parked_reader_test() { + let payload = encode(minimal_headers()) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let reply_to = process.new_subject() + let assert Ok(entry) = dict.get(state.streams, 1) + let state = + connection.State( + ..state, + streams: dict.insert( + state.streams, + 1, + connection.Stream(..entry, parked_reader: Some(reply_to)), + ), + ) + + let trailer_payload = encode([custom_field()]) + let result = connection.handle_headers(state, 1, True, True, trailer_payload) + let assert connection.Proceed(_state) = result + + let assert Ok(http2.DoneEvent(trailers)) = process.receive(reply_to, 100) + assert trailers == [#("x-custom", "some-longer-header-value-here")] +} + +pub fn trailers_with_invalid_header_is_protocol_error_test() { + let payload = encode(minimal_headers()) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let bad_trailer = + encode([ + alpacki.HeaderField( + <<"X-Bad":utf8>>, + <<"v":utf8>>, + alpacki.WithoutIndexing, + ), + ]) + let result = connection.handle_headers(state, 1, True, True, bad_trailer) + + let assert connection.RejectStream(_, 1, frame.ProtocolError) = result +} + +pub fn content_length_exceeded_rejects_stream_test() { + let payload = encode(headers_with_content_length("5")) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let result = connection.handle_data(state, 1, False, <<"toolong":utf8>>, 7) + + let assert connection.RejectStream(_, 1, frame.ProtocolError) = result +} + +pub fn content_length_short_at_end_stream_rejects_stream_test() { + let payload = encode(headers_with_content_length("5")) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let result = connection.handle_data(state, 1, True, <<"abc":utf8>>, 3) + + let assert connection.RejectStream(_, 1, frame.ProtocolError) = result +} + +pub fn content_length_matching_body_is_accepted_test() { + let payload = encode(headers_with_content_length("5")) + let assert connection.Proceed(state) = + connection.handle_headers(connection.test_state(), 1, False, True, payload) + + let result = connection.handle_data(state, 1, True, <<"abcde":utf8>>, 5) + + assert result != connection.RejectStream(state, 1, frame.ProtocolError) +} + +pub fn content_length_nonzero_with_no_body_rejects_before_spawn_test() { + let payload = encode(headers_with_content_length("5")) + + let result = + connection.handle_headers(connection.test_state(), 1, True, True, payload) + + let assert connection.RejectStream(state, 1, frame.ProtocolError) = result + assert dict.get(state.streams, 1) == Error(Nil) +} diff --git a/test/ewe/internal/http2/frame_test.gleam b/test/ewe/internal/http2/frame_test.gleam new file mode 100644 index 0000000..fba0037 --- /dev/null +++ b/test/ewe/internal/http2/frame_test.gleam @@ -0,0 +1,217 @@ +import ewe/internal/http2/frame +import gleam/bit_array + +pub fn data_frame_test() { + assert <<0:size(24), 0x0:8, 0x1:8, 0:1, 1:31, "":utf8>> + |> frame.decode(16_384) + == Ok(#(frame.Data(1, True, <<>>, 0), <<>>)) +} + +pub fn data_frame_with_payload_test() { + assert <<5:size(24), 0x0:8, 0x0:8, 0:1, 3:31, "hello":utf8, "extra":utf8>> + |> frame.decode(16_384) + == Ok(#(frame.Data(3, False, <<"hello":utf8>>, 5), <<"extra":utf8>>)) +} + +pub fn incomplete_header_test() { + assert frame.decode(<<0:size(24), 0x0:8>>, 16_384) == Error(frame.Incomplete) +} + +pub fn incomplete_payload_test() { + assert <<5:size(24), 0x0:8, 0x0:8, 0:1, 1:31, "hi":utf8>> + |> frame.decode(16_384) + == Error(frame.Incomplete) +} + +pub fn settings_frame_test() { + assert << + 12:size(24), 0x4:8, 0x0:8, 0:1, 0:31, 0x1:16, 4096:32, 0x3:16, 100:32, + >> + |> frame.decode(16_384) + == Ok( + #( + frame.Settings(0, False, [ + frame.HeaderTableSize(4096), + frame.MaxConcurrentStreams(100), + ]), + <<>>, + ), + ) +} + +pub fn settings_ack_test() { + assert <<0:size(24), 0x4:8, 0x1:8, 0:1, 0:31>> + |> frame.decode(16_384) + == Ok(#(frame.Settings(0, True, []), <<>>)) +} + +pub fn settings_ack_with_payload_is_frame_size_error_test() { + assert <<6:size(24), 0x4:8, 0x1:8, 0:1, 0:31, 0x1:16, 4096:32>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.FrameSizeError)) +} + +pub fn padded_data_frame_test() { + assert <<5:size(24), 0x0:8, 0x8:8, 0:1, 1:31, 2:8, "hi":utf8, 0:8, 0:8>> + |> frame.decode(16_384) + == Ok(#(frame.Data(1, False, <<"hi":utf8>>, 5), <<>>)) +} + +pub fn headers_with_priority_test() { + assert << + 10:size(24), 0x1:8, 0x24:8, 0:1, 1:31, 0:1, 0:31, 16:8, "block":utf8, + >> + |> frame.decode(16_384) + == Ok(#(frame.Headers(1, False, True, <<"block":utf8>>), <<>>)) +} + +pub fn window_update_test() { + assert <<4:size(24), 0x8:8, 0x0:8, 0:1, 5:31, 0:1, 1000:31>> + |> frame.decode(16_384) + == Ok(#(frame.WindowUpdate(5, 1000), <<>>)) +} + +pub fn unknown_frame_type_test() { + assert <<3:size(24), 0xf:8, 0x0:8, 0:1, 0:31, "abc":utf8>> + |> frame.decode(16_384) + == Ok(#(frame.Unknown(0, 0xf, <<"abc":utf8>>), <<>>)) +} + +pub fn rst_stream_error_code_test() { + assert <<4:size(24), 0x3:8, 0x0:8, 0:1, 1:31, 0x1:32>> + |> frame.decode(16_384) + == Ok(#(frame.RstStream(1, frame.ProtocolError), <<>>)) +} + +pub fn rst_stream_wrong_size_test() { + assert <<2:size(24), 0x3:8, 0x0:8, 0:1, 1:31, 0:16>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.FrameSizeError)) +} + +pub fn goaway_unknown_error_code_test() { + assert <<8:size(24), 0x7:8, 0x0:8, 0:1, 0:31, 0:1, 3:31, 999:32>> + |> frame.decode(16_384) + == Ok(#(frame.Goaway(0, 3, frame.UnknownErrorCode(999), <<>>), <<>>)) +} + +pub fn settings_enable_push_test() { + assert <<6:size(24), 0x4:8, 0x0:8, 0:1, 0:31, 0x2:16, 0:32>> + |> frame.decode(16_384) + == Ok(#(frame.Settings(0, False, [frame.EnablePush(False)]), <<>>)) +} + +pub fn settings_enable_push_invalid_value_test() { + assert <<6:size(24), 0x4:8, 0x0:8, 0:1, 0:31, 0x2:16, 2:32>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn settings_initial_window_size_overflow_test() { + assert <<6:size(24), 0x4:8, 0x0:8, 0:1, 0:31, 0x4:16, 2_147_483_648:32>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.FlowControlError)) +} + +pub fn settings_max_frame_size_out_of_range_test() { + assert <<6:size(24), 0x4:8, 0x0:8, 0:1, 0:31, 0x5:16, 100:32>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn padding_exceeds_payload_is_protocol_error_test() { + assert <<3:size(24), 0x0:8, 0x8:8, 0:1, 1:31, 10:8, "hi":utf8>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn frame_exceeding_max_frame_size_is_frame_size_error_test() { + assert <<5:size(24), 0x0:8, 0x0:8, 0:1, 3:31, "hello":utf8, "extra":utf8>> + |> frame.decode(4) + == Error(frame.Violation(frame.FrameSizeError)) +} + +pub fn priority_frame_with_stream_id_zero_is_protocol_error_test() { + assert <<5:size(24), 0x2:8, 0x0:8, 0:1, 0:31, 0:1, 2:31, 16:8>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn priority_frame_self_dependency_is_protocol_error_test() { + assert <<5:size(24), 0x2:8, 0x0:8, 0:1, 1:31, 0:1, 1:31, 16:8>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn priority_frame_test() { + assert <<5:size(24), 0x2:8, 0x0:8, 0:1, 1:31, 0:1, 2:31, 16:8>> + |> frame.decode(16_384) + == Ok(#(frame.Priority(1), <<>>)) +} + +pub fn rst_stream_with_stream_id_zero_is_protocol_error_test() { + assert <<4:size(24), 0x3:8, 0x0:8, 0:1, 0:31, 0x1:32>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn goaway_with_nonzero_stream_id_is_protocol_error_test() { + assert <<8:size(24), 0x7:8, 0x0:8, 0:1, 1:31, 0:1, 3:31, 999:32>> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn headers_with_priority_self_dependency_is_protocol_error_test() { + assert << + 10:size(24), 0x1:8, 0x24:8, 0:1, 1:31, 0:1, 1:31, 16:8, "block":utf8, + >> + |> frame.decode(16_384) + == Error(frame.Violation(frame.ProtocolError)) +} + +pub fn encode_data_roundtrip_test() { + let frame = frame.Data(1, True, <<"hello":utf8>>, 5) + assert frame.encode(frame) |> frame.decode(16_384) == Ok(#(frame, <<>>)) +} + +pub fn encode_headers_roundtrip_test() { + let frame = frame.Headers(3, False, True, <<"block":utf8>>) + assert frame.encode(frame) |> frame.decode(16_384) == Ok(#(frame, <<>>)) +} + +pub fn encode_settings_roundtrip_test() { + let frame = + frame.Settings(0, False, [ + frame.MaxConcurrentStreams(100), + frame.EnablePush(False), + ]) + assert frame.encode(frame) |> frame.decode(16_384) == Ok(#(frame, <<>>)) +} + +pub fn encode_rst_stream_roundtrip_test() { + let frame = frame.RstStream(5, frame.Cancel) + assert frame.encode(frame) |> frame.decode(16_384) == Ok(#(frame, <<>>)) +} + +pub fn encode_goaway_roundtrip_test() { + let frame = frame.Goaway(0, 7, frame.EnhanceYourCalm, <<"bye":utf8>>) + assert frame.encode(frame) |> frame.decode(16_384) == Ok(#(frame, <<>>)) +} + +pub fn encode_window_update_roundtrip_test() { + let frame = frame.WindowUpdate(9, 65_535) + assert frame.encode(frame) |> frame.decode(16_384) == Ok(#(frame, <<>>)) +} + +pub fn encode_ping_roundtrip_test() { + let frame = frame.Ping(0, True, <<1, 2, 3, 4, 5, 6, 7, 8>>) + assert frame.encode(frame) |> frame.decode(16_384) == Ok(#(frame, <<>>)) +} + +pub fn encode_data_header_declares_length_without_payload_test() { + let header = frame.encode_data_header(1, True, 5) + assert header + |> bit_array.append(<<"hello":utf8>>) + |> frame.decode(16_384) + == Ok(#(frame.Data(1, True, <<"hello":utf8>>, 5), <<>>)) +} -- 2.51.2