diff --git a/README.md b/README.md index 65cb1f1..e11b267 100644 --- a/README.md +++ b/README.md @@ -5,6 +5,39 @@ ## A glistening Gleam web server. +### Bounded WebSocket connections + +Use `mist.websocket_with_options` when clients must not control how much +message data a connection retains. Its callbacks match `mist.websocket`: + +```gleam +mist.websocket_with_options( + request: req, + on_init: fn(_connection) { #(Nil, None) }, + on_close: fn(_state) { Nil }, + handler: handle_ws_message, + options: mist.WebsocketOptions( + max_frame_bytes: 65_536, + max_message_bytes: 262_144, + compression: mist.CompressionDisabled, + ), +) +``` + +Both limits must be positive. The frame limit applies to its declared payload; +the message limit includes all continuation fragments, even across TCP reads. +Oversized input closes with code 1009 before the application receives it. +Detected framing violations or compressed frames close with code 1002. Bounded +connections do not negotiate compression and reject compressed frames before +decompression. Other protocol validation retains the existing decoder's behavior. + +These are per-connection input limits. Applications must separately bound +connection count, queued output, and data retained by their callbacks. A socket +send timeout does not mean that a remote application consumed the output. +The existing `mist.websocket` API retains its previous behavior. + +### Example application + To follow along with the example below, you can create a new project and add the dependencies as follows: diff --git a/gleam.toml b/gleam.toml index 821afff..becadf5 100644 --- a/gleam.toml +++ b/gleam.toml @@ -18,7 +18,7 @@ hpack_erl = ">= 0.3.0 and < 1.0.0" logging = ">= 1.0.0 and < 2.0.0" glisten = ">= 9.0.0 and < 10.0.0" exception = ">= 2.1.0 and < 3.0.0" -gramps = ">= 6.0.0 and < 7.0.0" +gramps = { git = "https://github.com/Roasbeef/gramps.git", ref = "a37a8ae3fe2531375b49b865c155eb2987d4d22a" } gleam_otp = ">= 1.2.0 and < 2.0.0" [dev-dependencies] diff --git a/manifest.toml b/manifest.toml index a605cee..e4d42e8 100644 --- a/manifest.toml +++ b/manifest.toml @@ -1,5 +1,10 @@ -# This file was generated by Gleam -# You typically do not need to edit this file +# Do not manually edit this file, it is managed by Gleam. +# +# This file locks the dependency versions used, to make your build +# deterministic and to prevent unexpected versions from being included +# in your application. +# +# You should check this file into your source control repository. packages = [ { name = "certifi", version = "2.15.0", build_tools = ["rebar3"], requirements = [], otp_app = "certifi", source = "hex", outer_checksum = "B147ED22CE71D72EAFDAD94F055165C1C182F61A2FF49DF28BCC71D1D5B94A60" }, @@ -12,7 +17,7 @@ packages = [ { name = "gleam_stdlib", version = "1.0.0", build_tools = ["gleam"], requirements = [], otp_app = "gleam_stdlib", source = "hex", outer_checksum = "960090C2FB391784BB34267B099DC9315CC1B1F6013E7415BC763CEF1905D7D3" }, { name = "gleeunit", version = "1.10.0", build_tools = ["gleam"], requirements = ["gleam_stdlib"], otp_app = "gleeunit", source = "hex", outer_checksum = "254B697FE72EEAD7BF82E941723918E421317813AC49923EE76A18C788C61E72" }, { name = "glisten", version = "9.0.1", build_tools = ["gleam"], requirements = ["gleam_erlang", "gleam_otp", "gleam_stdlib", "logging"], otp_app = "glisten", source = "hex", outer_checksum = "7795AA50830656F3A0316A6B26595F893C83272DA901B3405E31339CAA31A10B" }, - { name = "gramps", version = "6.0.1", build_tools = ["gleam"], requirements = ["gleam_crypto", "gleam_erlang", "gleam_http", "gleam_stdlib"], otp_app = "gramps", source = "hex", outer_checksum = "D55636072DEE173F6586A5679D3C02EC7A0DE3F8646B78C351B72908FF223DF7" }, + { name = "gramps", version = "6.0.1", build_tools = ["gleam"], requirements = ["gleam_crypto", "gleam_erlang", "gleam_http", "gleam_stdlib"], source = "git", repo = "https://github.com/Roasbeef/gramps.git", commit = "a37a8ae3fe2531375b49b865c155eb2987d4d22a" }, { name = "hackney", version = "1.25.0", build_tools = ["rebar3"], requirements = ["certifi", "idna", "metrics", "mimerl", "parse_trans", "ssl_verify_fun", "unicode_util_compat"], otp_app = "hackney", source = "hex", outer_checksum = "7209BFD75FD1F42467211FF8F59EA74D6F2A9E81CBCEE95A56711EE79FD6B1D4" }, { name = "hpack_erl", version = "0.3.0", build_tools = ["rebar3"], requirements = [], otp_app = "hpack", source = "hex", outer_checksum = "D6137D7079169D8C485C6962DFE261AF5B9EF60FBC557344511C1E65E3D95FB0" }, { name = "idna", version = "6.1.1", build_tools = ["rebar3"], requirements = ["unicode_util_compat"], otp_app = "idna", source = "hex", outer_checksum = "92376EB7894412ED19AC475E4A86F7B413C1B9FBB5BD16DCCD57934157944CEA" }, @@ -33,6 +38,6 @@ gleam_otp = { version = ">= 1.2.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 = { version = ">= 9.0.0 and < 10.0.0" } -gramps = { version = ">= 6.0.0 and < 7.0.0" } +gramps = { git = "https://github.com/Roasbeef/gramps.git", ref = "a37a8ae3fe2531375b49b865c155eb2987d4d22a" } hpack_erl = { version = ">= 0.3.0 and < 1.0.0" } logging = { version = ">= 1.0.0 and < 2.0.0" } diff --git a/src/mist.gleam b/src/mist.gleam index b10e3fd..b90bf24 100644 --- a/src/mist.gleam +++ b/src/mist.gleam @@ -19,6 +19,7 @@ import gleam/string_tree.{type StringTree} import glisten import glisten/transport import gramps/websocket.{BinaryFrame, Data, TextFrame} as gramps_websocket +import gramps/websocket/decoder as websocket_decoder import logging import mist/internal/buffer.{type Buffer, Buffer} import mist/internal/encoder @@ -629,6 +630,74 @@ pub fn websocket( on_init on_init: fn(WebsocketConnection) -> #(state, Option(process.Selector(message))), on_close on_close: fn(state) -> Nil, +) -> Response(ResponseData) { + let extensions = + request + |> request.get_header("sec-websocket-extensions") + |> result.map(fn(header) { string.split(header, ";") }) + |> result.unwrap([]) + websocket_upgrade(request, handler, on_init, on_close, extensions, None) +} + +/// Compression policy for bounded WebSocket connections. Compression stays +/// disabled until the decoder can bound decompressed output as well as input. +pub type WebsocketCompression { + /// Do not negotiate permessage-deflate; compressed frames are refused. + CompressionDisabled +} + +/// Per-connection payload limits, independent from application authorization +/// and total server connection or output queue budgets. +pub type WebsocketOptions { + WebsocketOptions( + /// Maximum payload bytes declared by one frame; must be positive. + max_frame_bytes: Int, + /// Maximum bytes across all fragments of one message; must be positive. + max_message_bytes: Int, + /// Compression cannot bypass the decoded-message limit. + compression: WebsocketCompression, + ) +} + +/// Upgrades with bounded, incremental decoding before application callbacks. +/// Invalid limits return HTTP 400 before the upgrade. Oversized payloads close +/// with code 1009; detected framing violations or compressed frames close with +/// code 1002. Other protocol validation retains the existing decoder's behavior. +/// +/// ## Examples +/// +/// ```gleam +/// let limits = WebsocketOptions(65_536, 65_536, CompressionDisabled) +/// // websocket_with_options(request, handler, on_init, on_close, limits) +/// ``` +pub fn websocket_with_options( + request request: Request(Connection), + handler handler: fn(state, WebsocketMessage(message), WebsocketConnection) -> + Next(state, message), + on_init on_init: fn(WebsocketConnection) -> + #(state, Option(process.Selector(message))), + on_close on_close: fn(state) -> Nil, + options options: WebsocketOptions, +) -> Response(ResponseData) { + case + websocket_decoder.new(options.max_frame_bytes, options.max_message_bytes) + { + Ok(decoder) -> + websocket_upgrade(request, handler, on_init, on_close, [], Some(decoder)) + Error(Nil) -> + response.new(400) |> response.set_body(Bytes(bytes_tree.new())) + } +} + +fn websocket_upgrade( + request: Request(Connection), + handler: fn(state, WebsocketMessage(message), WebsocketConnection) -> + Next(state, message), + on_init: fn(WebsocketConnection) -> + #(state, Option(process.Selector(message))), + on_close: fn(state) -> Nil, + extensions: List(String), + decoder: Option(websocket_decoder.Decoder), ) -> Response(ResponseData) { let handler = fn(state, message, connection) { message @@ -637,24 +706,19 @@ pub fn websocket( |> result.unwrap(continue(state)) |> convert_next } - let extensions = - request - |> request.get_header("sec-websocket-extensions") - |> result.map(fn(header) { string.split(header, ";") }) - |> result.unwrap([]) - let socket = request.body.socket let transport = request.body.transport case http.upgrade(socket, transport, extensions, request) { Ok(_nil) -> { let start = fn() { - websocket.initialize_connection( + websocket.initialize_connection_with_decoder( on_init, on_close, handler, socket, transport, extensions, + decoder, ) } diff --git a/src/mist/internal/websocket.gleam b/src/mist/internal/websocket.gleam index bc5ba0e..633261a 100644 --- a/src/mist/internal/websocket.gleam +++ b/src/mist/internal/websocket.gleam @@ -11,6 +11,7 @@ import glisten/socket/options import glisten/transport.{type Transport} import gramps/websocket.{type Frame, CloseFrame, Control, PingFrame} import gramps/websocket/compression.{type Compression, type Context} +import gramps/websocket/decoder import logging import mist/internal/next.{type Next, AbnormalStop, Continue, NormalStop} @@ -43,6 +44,7 @@ pub type WebsocketState(state) { buffer: BitArray, user: state, permessage_deflate: Option(Compression), + decoder: Option(decoder.Decoder), ) } @@ -88,6 +90,27 @@ pub fn initialize_connection( socket: Socket, transport: Transport, extensions: List(String), +) -> Result(actor.Started(process.Pid), actor.StartError) { + initialize_connection_with_decoder( + on_init, + on_close, + handler, + socket, + transport, + extensions, + None, + ) +} + +/// Starts a connection with an optional bounded, uncompressed decoder. +pub fn initialize_connection_with_decoder( + on_init: fn(WebsocketConnection) -> #(state, Option(Selector(user_message))), + on_close: fn(state) -> Nil, + handler: Handler(state, user_message), + socket: Socket, + transport: Transport, + extensions: List(String), + decoder: Option(decoder.Decoder), ) -> Result(actor.Started(process.Pid), actor.StartError) { let takeovers = websocket.get_context_takeovers(extensions) actor.new_with_initialiser(500, fn(subject) { @@ -114,6 +137,7 @@ pub fn initialize_connection( buffer: <<>>, user: initial_state, permessage_deflate: compression, + decoder:, ) |> actor.initialised |> actor.selecting(selector) @@ -131,30 +155,59 @@ pub fn initialize_connection( ) case msg { Valid(SocketMessage(data)) -> { - let #(frames, rest) = - websocket.decode_many_frames( - <>, - option.map(state.permessage_deflate, fn(compression) { - compression.inflate - }), - [], - ) - frames - |> websocket.aggregate_frames(None, []) - |> result.map(fn(frames) { - let next = - apply_frames( - frames, - handler, - connection, - Continue(state.user, None), - on_close, - ) + let decoded = case state.decoder { + Some(bounded) -> { + let #(bounded, next) = + receive_bounded( + bounded, + data, + handler, + connection, + Continue(state.user, None), + on_close, + ) + Ok(#(<<>>, Some(bounded), next)) + } + None -> { + let #(frames, rest) = + websocket.decode_many_frames( + <>, + option.map(state.permessage_deflate, fn(compression) { + compression.inflate + }), + [], + ) + frames + |> websocket.aggregate_frames(None, []) + |> result.map(fn(frames) { + let next = + apply_frames( + frames, + handler, + connection, + Continue(state.user, None), + on_close, + ) + #(rest, None, next) + }) + } + } + decoded + |> result.map(fn(decoded) { + let #(rest, bounded, next) = decoded case next { Continue(user_state, selector) -> { + // Rearm once after the entire TCP chunk has been handled. Doing + // this per decoded frame admits extra chunks into the mailbox. + set_active(connection.transport, connection.socket) let next = actor.continue( - WebsocketState(..state, buffer: rest, user: user_state), + WebsocketState( + ..state, + buffer: rest, + user: user_state, + decoder: bounded, + ), ) case selector { Some(selector) -> actor.with_selector(next, selector) @@ -288,6 +341,51 @@ pub fn initialize_connection( }) } +// Deliver one message at a time, leaving later input undecoded while its user +// callback runs. The decoder retains fragments across TCP chunks and rejects +// an oversized declared payload before copying or unmasking its contents. +fn receive_bounded( + bounded: decoder.Decoder, + data: BitArray, + handler: Handler(state, user_message), + connection: WebsocketConnection, + next: Next(state, WebsocketMessage(user_message)), + on_close: fn(state) -> Nil, +) -> #(decoder.Decoder, Next(state, WebsocketMessage(user_message))) { + case next { + NormalStop | AbnormalStop(_) -> #(bounded, next) + Continue(user, _) -> + case decoder.next(bounded, data) { + Ok(decoder.More(bounded)) -> #(bounded, next) + Ok(decoder.Frame(frame, bounded, rest)) -> { + let next = apply_frames([frame], handler, connection, next, on_close) + receive_bounded(bounded, rest, handler, connection, next, on_close) + } + Error(error) -> { + let reason = case error { + decoder.FrameTooLarge | decoder.MessageTooLarge -> + websocket.MessageTooBig(<<>>) + decoder.CompressionUnsupported | decoder.InvalidFrame -> + websocket.ProtocolError(<<>>) + } + let _ = + transport.send( + connection.transport, + connection.socket, + websocket.encode_close_frame(reason, None), + ) + on_close(user) + #( + bounded, + AbnormalStop( + "WebSocket input exceeded its limits or violated the protocol", + ), + ) + } + } + } +} + fn apply_frames( frames: List(Frame), handler: Handler(state, user_message), @@ -298,10 +396,7 @@ fn apply_frames( case frames, next { _, AbnormalStop(reason) -> AbnormalStop(reason) _, NormalStop -> NormalStop - [], next -> { - set_active(connection.transport, connection.socket) - next - } + [], next -> next [Control(CloseFrame(reason)), ..], Continue(state, _selector) -> { let _ = transport.send( @@ -321,7 +416,6 @@ fn apply_frames( websocket.encode_pong_frame(payload, None), ) |> result.map(fn(_nil) { - set_active(connection.transport, connection.socket) apply_frames(rest, handler, connection, continue, on_close) }) |> result.lazy_unwrap(fn() { diff --git a/test/bounded_websocket_test.gleam b/test/bounded_websocket_test.gleam new file mode 100644 index 0000000..cfa830a --- /dev/null +++ b/test/bounded_websocket_test.gleam @@ -0,0 +1,166 @@ +import exception +import gleam/bit_array +import gleam/bytes_tree +import gleam/erlang/atom +import gleam/erlang/process +import gleam/option.{None} +import gleam/string +import glisten/socket.{type Socket, type SocketReason, Closed} +import glisten/socket/options +import glisten/tcp +import mist + +@external(erlang, "gen_tcp", "connect") +fn connect( + host: #(Int, Int, Int, Int), + port: Int, + options: List(options.ErlangTcpOption), + timeout: Int, +) -> Result(Socket, SocketReason) + +@external(erlang, "gen_server", "stop") +fn stop_server( + pid: process.Pid, + reason: process.ExitReason, + timeout: Int, +) -> atom.Atom + +fn with_socket(run: fn(Socket, process.Subject(String)) -> Nil) -> Nil { + let ports = process.new_subject() + let delivered = process.new_subject() + let assert Ok(server) = + mist.new(fn(request) { + mist.websocket_with_options( + request:, + options: mist.WebsocketOptions(8, 10, mist.CompressionDisabled), + on_init: fn(_) { #(Nil, None) }, + on_close: fn(_) { Nil }, + handler: fn(state, event, socket) { + case event { + mist.Text(text) -> { + process.send(delivered, text) + let assert Ok(Nil) = mist.send_text_frame(socket, text) + as "the bounded application echoes the complete message" + mist.continue(state) + } + mist.Binary(_) -> mist.continue(state) + mist.Closed | mist.Shutdown -> mist.stop() + mist.Custom(_) -> mist.continue(state) + } + }, + ) + }) + |> mist.bind("127.0.0.1") + |> mist.port(0) + |> mist.after_start(fn(port, _, _) { process.send(ports, port) }) + |> mist.start + as "the fixture listener starts on an ephemeral port" + let assert Ok(port) = process.receive(ports, 1000) + as "listener publishes its port" + let result = + exception.rescue(fn() { + let assert Ok(socket) = + connect( + #(127, 0, 0, 1), + port, + options.to_erl_options([ + options.Mode(options.Binary), + options.ActiveMode(options.Passive), + ]), + 1000, + ) + as "the raw client connects" + let handshake = + "GET / HTTP/1.1\r\nHost: localhost\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Extensions: permessage-deflate\r\n\r\n" + let assert Ok(Nil) = tcp.send(socket, bytes_tree.from_string(handshake)) + as "the client offers compression during upgrade" + let headers = read_headers(socket, "", 8) + assert string.contains(headers, "101 Switching Protocols") + assert !string.contains( + string.lowercase(headers), + "sec-websocket-extensions:", + ) + let outcome = exception.rescue(fn() { run(socket, delivered) }) + let _ = tcp.close(socket) + let assert Ok(Nil) = outcome as "the socket assertions pass" + Nil + }) + let _ = stop_server(server.pid, process.Normal, 1000) + let assert Ok(Nil) = result as "the bounded socket fixture completes" + Nil +} + +fn read_headers(socket: Socket, collected: String, remaining: Int) -> String { + case string.contains(collected, "\r\n\r\n"), remaining { + True, _ -> collected + False, 0 -> panic as "upgrade headers exceeded the receive budget" + False, _ -> { + let assert Ok(bytes) = tcp.receive_timeout(socket, 0, 1000) + as "the upgrade responds within its deadline" + let assert Ok(text) = bit_array.to_string(bytes) + as "HTTP headers are UTF-8" + let collected = collected <> text + assert string.byte_size(collected) <= 4096 + read_headers(socket, collected, remaining - 1) + } + } +} + +fn send(socket: Socket, bytes: BitArray) -> Nil { + let assert Ok(Nil) = tcp.send(socket, bytes_tree.from_bit_array(bytes)) + as "the test sends one frame or frame prefix" + Nil +} + +fn closed_with(socket: Socket, code: Int) -> Nil { + let assert Ok(<<0x88, 2, actual:16>>) = tcp.receive_timeout(socket, 4, 1000) + as "the server sends the expected close frame without waiting for payload" + assert actual == code + assert tcp.receive_timeout(socket, 1, 1000) == Error(Closed) + Nil +} + +pub fn oversized_declared_length_closes_before_payload_test() { + with_socket(fn(socket, delivered) { + send(socket, <<0x81, 0xff, 1_000_000_000:64>>) + closed_with(socket, 1009) + assert process.receive(delivered, 0) == Error(Nil) + }) +} + +pub fn fragmented_message_limit_survives_separate_socket_reads_test() { + with_socket(fn(socket, delivered) { + send(socket, <<0x01, 0x86, 0, 0, 0, 0, "123456">>) + + // A ping/pong is an ordering barrier: the previous fragment has been + // consumed before the final fragment arrives in a later socket read. + send(socket, <<0x89, 0x80, 0, 0, 0, 0>>) + assert tcp.receive_timeout(socket, 2, 1000) == Ok(<<0x8a, 0>>) + send(socket, <<0x80, 0x85>>) + closed_with(socket, 1009) + assert process.receive(delivered, 0) == Error(Nil) + }) +} + +pub fn compression_is_not_negotiated_and_rsv1_is_rejected_test() { + with_socket(fn(socket, delivered) { + send(socket, <<0xc1, 0x80>>) + closed_with(socket, 1002) + assert process.receive(delivered, 0) == Error(Nil) + }) +} + +pub fn normal_messages_and_fragments_reach_the_handler_test() { + with_socket(fn(socket, delivered) { + send(socket, <<0x81, 0x82, 0, 0, 0, 0, "ok">>) + assert tcp.receive_timeout(socket, 4, 1000) == Ok(<<0x81, 2, "ok">>) + assert process.receive(delivered, 1000) == Ok("ok") + send(socket, <<0x01, 0x86, 0, 0, 0, 0, "123456">>) + send(socket, <<0x89, 0x80, 0, 0, 0, 0>>) + assert tcp.receive_timeout(socket, 2, 1000) == Ok(<<0x8a, 0>>) + send(socket, <<0x80, 0x84, 0, 0, 0, 0, "7890">>) + assert tcp.receive_timeout(socket, 12, 1000) + == Ok(<<0x81, 10, "1234567890">>) + assert process.receive(delivered, 1000) == Ok("1234567890") + }) +}