From 360e966ab467d8334aea714a50c9c8a7bdba77b1 Mon Sep 17 00:00:00 2001 From: Halil Durak Date: Thu, 30 Jul 2026 14:21:46 +0300 Subject: [PATCH] rework `Connection.upgrade` with `header_parser` --- src/cdp/Connection.zig | 114 +++++++++++++++++++++------------- src/network/header_parser.zig | 66 ++++++++++++++++++++ 2 files changed, 138 insertions(+), 42 deletions(-) diff --git a/src/cdp/Connection.zig b/src/cdp/Connection.zig index d86ea0362..6a04a0569 100644 --- a/src/cdp/Connection.zig +++ b/src/cdp/Connection.zig @@ -28,6 +28,7 @@ const ArenaPool = @import("../ArenaPool.zig"); const WS = @import("../network/WS.zig"); const sys_net = @import("../sys/net.zig"); +const header_parser = @import("../network/header_parser.zig"); const log = lp.log; const posix = std.posix; @@ -418,59 +419,96 @@ fn pushCdp(self: *Connection, bytes: []const u8) !bool { } pub fn upgrade(self: *Connection, request: []u8) !void { - // our caller already confirmed that we have a trailing \r\n\r\n - const request_line_end = std.mem.indexOfScalar(u8, request, '\r') orelse unreachable; - const request_line = request[0..request_line_end]; - - if (!std.ascii.endsWithIgnoreCase(request_line, "http/1.1")) { + var cursor = header_parser.Cursor{ + .idx = request.ptr, + .start = request.ptr, + .end = request.ptr + request.len, + }; + // A malformed request line maps to a 400 in processHttpRequest. + header_parser.validateWebSocketRequestLine(&cursor) catch { return error.InvalidProtocol; - } + }; - // we need to extract the sec-websocket-key value - var key: []const u8 = ""; + // We need to extract the `Sec-WebSocket-Key` value. + var sec_websocket_key: []const u8 = ""; + // We need to make sure that we got all the necessary headers + values. + const RequiredHeaders = packed struct(u8) { + /// Upgrade: websocket + upgrade: bool = false, + sec_websocket_version: bool = false, + /// Connection: upgrade + connection: bool = false, + sec_websocket_key: bool = false, + __pad: u4 = 0, + }; - // we need to make sure that we got all the necessary headers + values - var required_headers: u8 = 0; + var required_headers = RequiredHeaders{}; + // We reuse this to parse headers that're required. + var header: header_parser.Header = undefined; + while (true) { + if (cursor.reachedEnd()) { + return error.InvalidRequest; + } - // can't std.mem.split because it forces the iterated value to be const - // (we could @constCast...) + // Check if headers part has finished. + switch (cursor.char()) { + '\n' => { + // End of headers. + cursor.advance(1); + break; + }, + '\r' => { + // We need an LF too. + if (!cursor.hasLength(2) or !cursor.peek2('\r', '\n')) { + return error.InvalidRequest; + } + // End of headers. + cursor.advance(2); + break; + }, + else => {}, + } - var buf = request[request_line_end + 2 ..]; + // A malformed header maps to a 400 in processHttpRequest. + header.parse(&cursor) catch { + return error.InvalidRequest; + }; + const key = header.key; + const value = header.value; - while (buf.len > 4) { - const index = std.mem.indexOfScalar(u8, buf, '\r') orelse unreachable; - const separator = std.mem.indexOfScalar(u8, buf[0..index], ':') orelse return error.InvalidRequest; - - const name = std.mem.trim(u8, toLower(buf[0..separator]), &std.ascii.whitespace); - const value = std.mem.trim(u8, buf[(separator + 1)..index], &std.ascii.whitespace); - - if (std.mem.eql(u8, name, "upgrade")) { + // Header names are case-insensitive; `Header.parse` keeps their + // original casing. + if (std.ascii.eqlIgnoreCase(key, "upgrade")) { if (!std.ascii.eqlIgnoreCase("websocket", value)) { return error.InvalidUpgradeHeader; } - required_headers |= 1; - } else if (std.mem.eql(u8, name, "sec-websocket-version")) { + required_headers.upgrade = true; + } else if (std.ascii.eqlIgnoreCase(key, "sec-websocket-version")) { if (value.len != 2 or value[0] != '1' or value[1] != '3') { return error.InvalidVersionHeader; } - required_headers |= 2; - } else if (std.mem.eql(u8, name, "connection")) { + required_headers.sec_websocket_version = true; + } else if (std.ascii.eqlIgnoreCase(key, "connection")) { // find if connection header has upgrade in it, example header: // Connection: keep-alive, Upgrade if (std.ascii.indexOfIgnoreCase(value, "upgrade") == null) { return error.InvalidConnectionHeader; } - required_headers |= 4; - } else if (std.mem.eql(u8, name, "sec-websocket-key")) { - key = value; - required_headers |= 8; + required_headers.connection = true; + } else if (std.ascii.eqlIgnoreCase(key, "sec-websocket-key")) { + sec_websocket_key = value; + required_headers.sec_websocket_key = true; } - - const next = index + 2; - buf = buf[next..]; } - if (required_headers != 15) { + // Check if we've received all related headers. + const satisfied = @as(u8, @bitCast(required_headers)) == @as(u8, @bitCast(RequiredHeaders{ + .upgrade = true, + .sec_websocket_version = true, + .connection = true, + .sec_websocket_key = true, + })); + if (!satisfied) { return error.MissingHeaders; } @@ -498,7 +536,7 @@ pub fn upgrade(self: *Connection, request: []u8) !void { const key_pos = res.len - 32; var h: [20]u8 = undefined; var hasher = std.crypto.hash.Sha1.init(.{}); - hasher.update(key); + hasher.update(sec_websocket_key); // websocket spec always used this value hasher.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); hasher.final(&h); @@ -583,11 +621,3 @@ fn websocketHeader(buf: []u8, op_code: WS.OpCode, payload_len: usize) []const u8 buf[9] = @intCast(len & 0xFF); return buf[0..10]; } - -// In-place string lowercase -fn toLower(str: []u8) []u8 { - for (str, 0..) |ch, i| { - str[i] = std.ascii.toLower(ch); - } - return str; -} diff --git a/src/network/header_parser.zig b/src/network/header_parser.zig index 54e672d03..522a6a7e3 100644 --- a/src/network/header_parser.zig +++ b/src/network/header_parser.zig @@ -231,6 +231,66 @@ pub const Header = struct { } }; +/// Validates WebSocket initialization requests. +/// Currently does not validate paths. +pub fn validateWebSocketRequestLine(cursor: *Cursor) !void { + // GET / HTTP/1.1\n + const min_request_len = 0xf; + if (cursor.hasLength(min_request_len) == false) { + return error.Incomplete; + } + + // WS requests can only be send w/ GET method. + if (!cursor.peek4('G', 'E', 'T', ' ')) { + return error.Invalid; + } + cursor.advance(4); + + const path_start = cursor.current(); + // Find the first space. + while (cursor.end - cursor.current() > 0 and cursor.char() != ' ') : (cursor.advance(1)) {} + const path_end = cursor.current(); + // 0 length path. + if (path_start == path_end) { + return error.Invalid; + } + const path = path_start[0 .. path_end - path_start]; + _ = path; + + // Skip past the delimiting space(s); the scan above guarantees we're on + // a space or at the end, and recipients may parse on whitespace + // boundaries (RFC 9112 ยง3). + while (cursor.end - cursor.current() > 0 and cursor.char() == ' ') : (cursor.advance(1)) {} + + // HTTP/1.1(\r)\n + if (cursor.hasLength(9) == false) { + return error.Incomplete; + } + // Make sure we got HTTP/1.1. + if (cursor.asInteger(u64) != @as(u64, @bitCast(@as([]const u8, "HTTP/1.1")[0..8].*))) { + return error.Invalid; + } + cursor.advance(8); + + // Trailing (CR)LF. + switch (cursor.char()) { + '\n' => cursor.advance(1), + '\r' => { + // We need an LF too. + if (!cursor.hasLength(2)) { + return error.Incomplete; + } + if (!cursor.peek2('\r', '\n')) { + @branchHint(.unlikely); + return error.Invalid; + } + cursor.advance(2); + }, + // Any other character is invalid. + else => return error.Invalid, + } +} + pub const Disposition = struct { name: ?[]const u8 = null, filename: ?[]const u8 = null, @@ -343,6 +403,12 @@ pub const Cursor = struct { return cursor.asInteger(u16) == @as(u16, @bitCast([2]u8{ c0, c1 })); } + /// Peek the current and next 3 characters but don't advance. + /// SAFETY: This function doesn't check if out of bounds reachable. + pub fn peek4(cursor: *const Cursor, c0: u8, c1: u8, c2: u8, c3: u8) bool { + return cursor.asInteger(u32) == @as(u32, @bitCast([4]u8{ c0, c1, c2, c3 })); + } + /// Moves the cursor until no leading spaces there are. pub fn skipSpaces(cursor: *Cursor) void { while (cursor.end - cursor.current() > 0 and (cursor.char() == ' ' or cursor.char() == '\t')) : (cursor.advance(1)) {}