From c14617f9eac6f7693ea3f8671cb45e4006f42d86 Mon Sep 17 00:00:00 2001 From: Karl Seguin Date: Thu, 17 Jul 2025 09:08:49 +0800 Subject: [PATCH 1/5] update for latest zig (writegate) --- build.zig | 6 +- readme.md | 8 +-- src/proto.zig | 6 +- src/server/handshake.zig | 29 ++++------ src/server/server.zig | 95 ++++++++++++++++---------------- src/server/thread_pool.zig | 6 +- src/websocket.zig | 2 +- support/autobahn/client/main.zig | 2 +- test_runner.zig | 70 ++++++++++------------- 9 files changed, 103 insertions(+), 121 deletions(-) diff --git a/build.zig b/build.zig index 7719f4a..1982079 100644 --- a/build.zig +++ b/build.zig @@ -5,6 +5,8 @@ pub fn build(b: *std.Build) !void { const optimize = b.standardOptimizeOption(.{}); const websocket_module = b.addModule("websocket", .{ + .target = target, + .optimize = optimize, .root_source_file = b.path("src/websocket.zig"), }); @@ -17,9 +19,7 @@ pub fn build(b: *std.Build) !void { { // run tests const tests = b.addTest(.{ - .root_source_file = b.path("src/websocket.zig"), - .target = target, - .optimize = optimize, + .root_module = websocket_module, .test_runner = .{ .path = b.path("test_runner.zig"), .mode = .simple }, }); tests.linkLibC(); diff --git a/readme.md b/readme.md index ad25320..1c18fd3 100644 --- a/readme.md +++ b/readme.md @@ -192,14 +192,14 @@ The call to `init` includes a `*websocket.Conn`. It is expected that handlers wi `close` takes an optional value where you can specify the `code` and/or `reason`: `conn.close(.{.code = 4000, .reason = "bye bye"})` Refer to [RFC6455](https://datatracker.ietf.org/doc/html/rfc6455#section-7.4.1) for valid codes. The `reason` must be <= 123 bytes. ### Writer -It's possible to get a `std.io.Writer` from a `*Conn`. Because websocket messages are framed, the writter will buffer the message in memory and requires an explicit "flush". Buffering requires an allocator. +It's possible to get a `*std.Io.Writer` from a `*Conn`. Because websocket messages are framed, the writter will buffer the message in memory and requires an explicit "send". Buffering requires an allocator. ```zig // .text or .binary var wb = conn.writeBuffer(allocator, .text); defer wb.deinit(); -try std.fmt.format(wb.writer(), "it's over {d}!!!", .{9000}); -try wb.flush(); +try wb.interface.print("it's over {d}!!!", .{9000}); +try wb.send(); ``` Consider using the `clientMessage` overload which accepts an allocator. Not only is this allocator fast (it's a thread-local buffer than fallsback to an arena), but it also eliminates the need to call `deinit`: @@ -211,7 +211,7 @@ pub fn clientMessage(h: *Handler, allocator: Allocator, data: []const u8) !void var wb = conn.writeBuffer(allocator, .text); try std.fmt.format(wb.writer(), "it's over {d}!!!", .{9000}); - try wb.flush(); + try wb.send(); } ``` diff --git a/src/proto.zig b/src/proto.zig index 74aded5..a3c3eca 100644 --- a/src/proto.zig +++ b/src/proto.zig @@ -317,7 +317,7 @@ pub const Reader = struct { if (is_continuation) { if (self.fragment) |*f| { if (f.compressed) { - return . {more, .{.data = try self.decompress(try f.last(payload)), .type = f.type}}; + return .{ more, .{ .data = try self.decompress(try f.last(payload)), .type = f.type } }; } return .{ more, .{ .type = f.type, .data = try f.last(payload) } }; } @@ -330,7 +330,7 @@ pub const Reader = struct { } if (compressed) { - return . {more, .{.data = try self.decompress(payload), .type = message_type}}; + return .{ more, .{ .data = try self.decompress(payload), .type = message_type } }; } // just a normal single-fragment message (most common case) @@ -428,7 +428,7 @@ pub const Reader = struct { .buf = try provider.pool.acquireOrCreate(), }; } else { - writer = .{ + writer = .{ .pooled = false, .provider = provider, .buf = try provider.allocator.alloc(u8, @intFromFloat(@as(f64, @floatFromInt(compressed.len)) * 1.25)), diff --git a/src/server/handshake.zig b/src/server/handshake.zig index 3f72e5b..76d401b 100644 --- a/src/server/handshake.zig +++ b/src/server/handshake.zig @@ -69,7 +69,7 @@ pub const Handshake = struct { headers.add(name, value); switch (std.meta.stringToEnum(SpecialHeader, name) orelse .none) { .upgrade => { - if (!ascii.eqlIgnoreCase("websocket", value)) { + if (!ascii.eqlIgnoreCase("websocket", value)) { return error.InvalidUpgrade; } required_headers |= 1; @@ -127,8 +127,7 @@ pub const Handshake = struct { "HTTP/1.1 101 Switching Protocols\r\n" ++ "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ - "Sec-Websocket-Accept: " - ; + "Sec-Websocket-Accept: "; @memcpy(buf[0..HEADER.len], HEADER); var pos = HEADER.len; @@ -168,7 +167,7 @@ pub const Handshake = struct { } for (headers.keys[0..headers.len], headers.values[0..headers.len]) |k, v| { - pos += (try std.fmt.bufPrint(buf[pos..], "\r\n{s}: {s}", .{k, v})).len; + pos += (try std.fmt.bufPrint(buf[pos..], "\r\n{s}: {s}", .{ k, v })).len; } const end = pos + 4; @@ -373,7 +372,7 @@ pub const Pool = struct { max_res_headers: usize, states: []*Handshake.State, - pub fn init(allocator: Allocator, count: usize, buffer_size: usize, max_req_headers: usize, max_res_headers: usize) !*Pool { + pub fn init(allocator: Allocator, count: usize, buffer_size: usize, max_req_headers: usize, max_res_headers: usize) !*Pool { const states = try allocator.alloc(*Handshake.State, count); errdefer allocator.free(states); @@ -535,8 +534,7 @@ test "handshake: reply" { "HTTP/1.1 101 Switching Protocols\r\n" ++ "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ - "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n\r\n" - ; + "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n\r\n"; try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, null, &buf)); } @@ -547,8 +545,7 @@ test "handshake: reply" { "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Sec-WebSocket-Extensions: permessage-deflate\r\n\r\n" - ; + "Sec-WebSocket-Extensions: permessage-deflate\r\n\r\n"; try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{}, &buf)); } @@ -559,8 +556,7 @@ test "handshake: reply" { "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n\r\n" - ; + "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n\r\n"; try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{ .client_no_context_takeover = true, .server_no_context_takeover = true, @@ -576,8 +572,7 @@ test "handshake: reply" { "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Set-Cookie: Yummy!\r\n\r\n" - ; + "Set-Cookie: Yummy!\r\n\r\n"; try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, null, &buf)); } @@ -589,8 +584,7 @@ test "handshake: reply" { "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ "Sec-WebSocket-Extensions: permessage-deflate\r\n" ++ - "Set-Cookie: Yummy!\r\n\r\n" - ; + "Set-Cookie: Yummy!\r\n\r\n"; try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{}, &buf)); } @@ -602,8 +596,7 @@ test "handshake: reply" { "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n" ++ - "Set-Cookie: Yummy!\r\n\r\n" - ; + "Set-Cookie: Yummy!\r\n\r\n"; try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{ .client_no_context_takeover = true, .server_no_context_takeover = true, @@ -703,7 +696,7 @@ fn testPool(p: *Pool) void { var hs = p.acquire() catch unreachable; std.debug.assert(hs.buf[0] == 0); hs.buf[0] = 255; - std.time.sleep(random.uintAtMost(u32, 100000)); + std.Thread.sleep(random.uintAtMost(u32, 100000)); hs.buf[0] = 0; p.release(hs); } diff --git a/src/server/server.zig b/src/server/server.zig index 0e1256a..1c33e6e 100644 --- a/src/server/server.zig +++ b/src/server/server.zig @@ -382,7 +382,7 @@ pub fn Blocking(comptime H: type) type { } if (hc.handler != null) { // if we have a handler, the our handshake completed - try conn_manager.setupCompression(hc, compression); + try conn_manager.setupCompression(hc, compression); break; } if (timestamp() > deadline) { @@ -430,7 +430,7 @@ pub fn Blocking(comptime H: type) type { if (conn_manager.count() == 0) { return; } - std.time.sleep(std.time.ns_per_ms * 100); + std.Thread.sleep(std.time.ns_per_ms * 100); } } @@ -520,7 +520,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { var it = self.loop.wait(timeout) catch |err| { log.err("failed to wait on events: {}", .{err}); - std.time.sleep(std.time.ns_per_s); + std.Thread.sleep(std.time.ns_per_s); continue; }; @@ -530,7 +530,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { if (data == 0) { self.accept(listener, now) catch |err| { log.err("accept error: {}", .{err}); - std.time.sleep(std.time.ns_per_ms); + std.Thread.sleep(std.time.ns_per_ms); }; continue; } @@ -642,7 +642,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { var success = false; if (hc.handler == null) { success = self.dataForHandshake(hc) catch |err| blk: { - log.err("({}) error processing handshake: {}", .{ hc.conn.address, err }); + log.err("({any}) error processing handshake: {}", .{ hc.conn.address, err }); break :blk false; }; } else { @@ -681,7 +681,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { conn_manager.inactive(hc); } - try conn_manager.setupCompression(hc, compression); + try conn_manager.setupCompression(hc, compression); return true; } }; @@ -740,7 +740,7 @@ fn NonBlockingBase(comptime H: type, comptime MANAGE_HS: bool) type { pub fn dataAvailable(self: *Self, hc: *HandlerConn(H), thread_buf: []u8) bool { return self._dataAvailable(hc, thread_buf) catch |err| { - log.err("({}) error processing client message: {}", .{ hc.conn.address, err }); + log.err("({any}) error processing client message: {}", .{ hc.conn.address, err }); return false; }; } @@ -1382,7 +1382,7 @@ pub const Conn = struct { var fbs = std.io.fixedBufferStream(data); _ = try compressor.compress(fbs.reader()); try compressor.flush(); - payload = writer.items[0..writer.items.len - 4]; + payload = writer.items[0 .. writer.items.len - 4]; if (c.reset) { c.compressor = try Conn.Compression.Type.init(writer.writer(), .{}); @@ -1447,8 +1447,13 @@ pub const Conn = struct { pub fn writeBuffer(self: *Conn, allocator: Allocator, op_code: OpCode) Writer { return .{ .conn = self, + .buf = .empty, .op_code = op_code, - .buf = std.ArrayList(u8).init(allocator), + .allocator = allocator, + .interface = .{ + .vtable = &.{ .drain = Writer.drain }, + .buffer = &.{}, + }, }; } @@ -1461,38 +1466,37 @@ pub const Conn = struct { pub const Writer = struct { conn: *Conn, op_code: OpCode, - buf: std.ArrayList(u8), + allocator: Allocator, + buf: std.ArrayListUnmanaged(u8), + interface: std.io.Writer, pub const Error = Allocator.Error; - pub const IOWriter = std.io.Writer(*Writer, error{OutOfMemory}, Writer.write); pub fn deinit(self: *Writer) void { - self.buf.deinit(); - } - - pub fn writer(self: *Writer) IOWriter { - return .{ .context = self }; + self.buf.deinit(self.allocator); } - pub fn write(self: *Writer, data: []const u8) Allocator.Error!usize { - try self.buf.appendSlice(data); - return data.len; + pub fn drain(io_w: *std.io.Writer, data: []const []const u8, splat: usize) error{WriteFailed}!usize { + _ = splat; + const self: *Writer = @alignCast(@fieldParentPtr("interface", io_w)); + self.buf.appendSlice(self.allocator, data[0]) catch return error.WriteFailed; + return data[0].len; } - pub fn flush(self: *Writer) !void { - try self.conn.writeFrame(self.op_code, self.buf.items); + pub fn send(self: *Writer) !void { + return self.conn.writeFrame(self.op_code, self.buf.items) catch error.WriteFailed; } }; }; -fn handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) struct{?Compression, bool} { +fn handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) struct { ?Compression, bool } { return _handleHandshake(H, worker, hc, ctx) catch |err| { - log.warn("({}) uncaugh error processing handshake: {}", .{ hc.conn.address, err }); - return .{null, false}; + log.warn("({any}) uncaugh error processing handshake: {}", .{ hc.conn.address, err }); + return .{ null, false }; }; } -fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) !struct{?Compression, bool} { +fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) !struct { ?Compression, bool } { std.debug.assert(hc.handler == null); var state = hc.handshake orelse blk: { @@ -1506,35 +1510,35 @@ fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: const len = state.len; if (len == buf.len) { - log.warn("({}) handshake request exceeded maximum configured size ({d})", .{ conn.address, buf.len }); - return .{null, false}; + log.warn("({any}) handshake request exceeded maximum configured size ({d})", .{ conn.address, buf.len }); + return .{ null, false }; } const n = posix.read(hc.socket, buf[len..]) catch |err| { switch (err) { - error.BrokenPipe, error.ConnectionResetByPeer => log.debug("({}) handshake connection closed: {}", .{ conn.address, err }), + error.BrokenPipe, error.ConnectionResetByPeer => log.debug("({any}) handshake connection closed: {}", .{ conn.address, err }), error.WouldBlock => { std.debug.assert(blockingMode()); - log.debug("({}) handshake timeout", .{conn.address}); + log.debug("({any}) handshake timeout", .{conn.address}); }, - else => log.warn("({}) handshake error reading from socket: {}", .{ conn.address, err }), + else => log.warn("({any}) handshake error reading from socket: {}", .{ conn.address, err }), } - return .{null, false}; + return .{ null, false }; }; if (n == 0) { - log.debug("({}) handshake connection closed", .{conn.address}); - return .{null, false}; + log.debug("({any}) handshake connection closed", .{conn.address}); + return .{ null, false }; } state.len = len + n; var handshake = Handshake.parse(state) catch |err| { - log.debug("({}) error parsing handshake: {}", .{ conn.address, err }); + log.debug("({any}) error parsing handshake: {}", .{ conn.address, err }); respondToHandshakeError(conn, err); - return .{null, false}; + return .{ null, false }; } orelse { // we need more data - return .{null, true}; + return .{ null, true }; }; var agreed_compression: ?Compression = null; @@ -1550,7 +1554,6 @@ fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: defer state.release(); hc.handshake = null; - // After this, the app has access to &hc.conn, so any access to the // conn has to be synchronized (which the conn does internally). @@ -1561,7 +1564,7 @@ fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: respondToHandshakeError(conn, err); } log.debug("({}) " ++ @typeName(H) ++ ".init rejected request {}", .{ conn.address, err }); - return .{null, false}; + return .{ null, false }; }; hc.handler = handler; @@ -1575,18 +1578,18 @@ fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: const res = if (params.len == 1) hc.handler.?.afterInit() else hc.handler.?.afterInit(ctx); res catch |err| { log.debug("({}) " ++ @typeName(H) ++ ".afterInit error: {}", .{ conn.address, err }); - return .{null, false}; + return .{ null, false }; }; } log.debug("({}) connection successfully upgraded", .{conn.address}); - return .{agreed_compression, true}; + return .{ agreed_compression, true }; } fn handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator, fba: *FixedBufferAllocator) bool { std.debug.assert(hc.handshake == null); return _handleClientData(H, hc, allocator, fba) catch |err| { - log.warn("({}) uncaugh error handling incoming data: {}", .{ hc.conn.address, err }); + log.warn("({any}) uncaugh error handling incoming data: {}", .{ hc.conn.address, err }); return false; }; } @@ -1597,7 +1600,7 @@ fn _handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator reader.fill(conn.stream) catch |err| { switch (err) { error.BrokenPipe, error.Closed, error.ConnectionResetByPeer => log.debug("({}) connection closed: {}", .{ conn.address, err }), - else => log.warn("({}) error reading from connection: {}", .{ conn.address, err }), + else => log.warn("({any}) error reading from connection: {}", .{ conn.address, err }), } return false; }; @@ -1612,7 +1615,7 @@ fn _handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator error.CompressionError => conn.writeFramed(CLOSE_PROTOCOL_ERROR) catch {}, else => {}, } - log.debug("({}) invalid websocket packet: {}", .{ conn.address, err }); + log.debug("({any}) invalid websocket packet: {}", .{ conn.address, err }); return false; } orelse { // everything is fine, we just need more data @@ -1622,7 +1625,7 @@ fn _handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator const message_type = message.type; defer reader.done(message_type); - log.debug("({}) received {s} message", .{ hc.conn.address, @tagName(message_type) }); + log.debug("({anys}) received {s} message", .{ hc.conn.address, @tagName(message_type) }); switch (message_type) { .text, .binary => { const params = @typeInfo(@TypeOf(H.clientMessage)).@"fn".params; @@ -1998,8 +2001,8 @@ const TestHandler = struct { } if (std.mem.eql(u8, data, "writer")) { var wb = self.conn.writeBuffer(allocator, .text); - try std.fmt.format(wb.writer(), "{d}!!!", .{9000}); - return wb.flush(); + try wb.interface.print("{d}!!!", .{9000}); + return wb.send(); } if (std.mem.eql(u8, data, "ping")) { var buf = [_]u8{ 'a', '-', 'p', 'i', 'n', 'g' }; diff --git a/src/server/thread_pool.zig b/src/server/thread_pool.zig index e859346..1f7d72b 100644 --- a/src/server/thread_pool.zig +++ b/src/server/thread_pool.zig @@ -195,7 +195,7 @@ test "ThreadPool: small fuzz" { tp.spawn(.{1}); } while (tp.empty() == false) { - std.time.sleep(std.time.ns_per_ms); + std.Thread.sleep(std.time.ns_per_ms); } tp.deinit(); try t.expectEqual(50_000, testSum); @@ -209,7 +209,7 @@ test "ThreadPool: large fuzz" { tp.spawn(.{1}); } while (tp.empty() == false) { - std.time.sleep(std.time.ns_per_ms); + std.Thread.sleep(std.time.ns_per_ms); } tp.deinit(); try t.expectEqual(50_000, testSum); @@ -220,5 +220,5 @@ fn testIncr(c: u64, buf: []u8) void { std.debug.assert(buf.len == 512); _ = @atomicRmw(u64, &testSum, .Add, c, .monotonic); // let the threadpool queue get backed up - std.time.sleep(std.time.ns_per_us * 100); + std.Thread.sleep(std.time.ns_per_us * 100); } diff --git a/src/websocket.zig b/src/websocket.zig index 6834157..3a20ec3 100644 --- a/src/websocket.zig +++ b/src/websocket.zig @@ -46,7 +46,7 @@ const t = @import("t.zig"); test "frameText" { { const framed = frameText(""); - try t.expectString(&[_]u8{ 129, 0}, &framed); + try t.expectString(&[_]u8{ 129, 0 }, &framed); } { diff --git a/support/autobahn/client/main.zig b/support/autobahn/client/main.zig index 7ec65d1..44437ac 100644 --- a/support/autobahn/client/main.zig +++ b/support/autobahn/client/main.zig @@ -32,7 +32,7 @@ pub fn main() !void { }; // wait 5 seconds for autobanh server to be up - std.time.sleep(std.time.ns_per_s * 5); + std.Thread.sleep(std.time.ns_per_s * 5); var buffer_provider = try websocket.bufferProvider(allocator, .{ .count = 10, .size = 32768, .max = 20_000_000 }); defer buffer_provider.deinit(); diff --git a/test_runner.zig b/test_runner.zig index e9718ce..d9c5e62 100644 --- a/test_runner.zig +++ b/test_runner.zig @@ -1,9 +1,7 @@ // in your build.zig, you can specify a custom test runner: // const tests = b.addTest(.{ -// .target = target, -// .optimize = optimize, -// .test_runner = .{ .path = b.path("test_runner.zig"), .mode = .simple }, // add this line -// .root_source_file = b.path("src/main.zig"), +// .root_module = $MODULE_BEING_TESTED, +// .test_runner = .{ .path = b.path("test_runner.zig"), .mode = .simple }, // }); pub const std_options = std.Options{ .log_scope_levels = &[_]std.log.ScopeLevel{ @@ -37,13 +35,12 @@ pub fn main() !void { var skip: usize = 0; var leak: usize = 0; - const printer = Printer.init(); - printer.fmt("\r\x1b[0K", .{}); // beginning of line and clear to end of line + Printer.fmt("\r\x1b[0K", .{}); // beginning of line and clear to end of line for (builtin.test_functions) |t| { if (isSetup(t)) { t.func() catch |err| { - printer.status(.fail, "\nsetup \"{s}\" failed: {}\n", .{ t.name, err }); + Printer.status(.fail, "\nsetup \"{s}\" failed: {}\n", .{ t.name, err }); return err; }; } @@ -85,7 +82,7 @@ pub fn main() !void { if (std.testing.allocator_instance.deinit() == .leak) { leak += 1; - printer.status(.fail, "\n{s}\n\"{s}\" - Memory Leak\n{s}\n", .{ BORDER, friendly_name, BORDER }); + Printer.status(.fail, "\n{s}\n\"{s}\" - Memory Leak\n{s}\n", .{ BORDER, friendly_name, BORDER }); } if (result) |_| { @@ -98,7 +95,7 @@ pub fn main() !void { else => { status = .fail; fail += 1; - printer.status(.fail, "\n{s}\n\"{s}\" - {s}\n{s}\n", .{ BORDER, friendly_name, @errorName(err), BORDER }); + Printer.status(.fail, "\n{s}\n\"{s}\" - {s}\n{s}\n", .{ BORDER, friendly_name, @errorName(err), BORDER }); if (@errorReturnTrace()) |trace| { std.debug.dumpStackTrace(trace.*); } @@ -110,16 +107,16 @@ pub fn main() !void { if (env.verbose) { const ms = @as(f64, @floatFromInt(ns_taken)) / 1_000_000.0; - printer.status(status, "{s} ({d:.2}ms)\n", .{ friendly_name, ms }); + Printer.status(status, "{s} ({d:.2}ms)\n", .{ friendly_name, ms }); } else { - printer.status(status, ".", .{}); + Printer.status(status, ".", .{}); } } for (builtin.test_functions) |t| { if (isTeardown(t)) { t.func() catch |err| { - printer.status(.fail, "\nteardown \"{s}\" failed: {}\n", .{ t.name, err }); + Printer.status(.fail, "\nteardown \"{s}\" failed: {}\n", .{ t.name, err }); return err; }; } @@ -127,43 +124,32 @@ pub fn main() !void { const total_tests = pass + fail; const status = if (fail == 0) Status.pass else Status.fail; - printer.status(status, "\n{d} of {d} test{s} passed\n", .{ pass, total_tests, if (total_tests != 1) "s" else "" }); + Printer.status(status, "\n{d} of {d} test{s} passed\n", .{ pass, total_tests, if (total_tests != 1) "s" else "" }); if (skip > 0) { - printer.status(.skip, "{d} test{s} skipped\n", .{ skip, if (skip != 1) "s" else "" }); + Printer.status(.skip, "{d} test{s} skipped\n", .{ skip, if (skip != 1) "s" else "" }); } if (leak > 0) { - printer.status(.fail, "{d} test{s} leaked\n", .{ leak, if (leak != 1) "s" else "" }); + Printer.status(.fail, "{d} test{s} leaked\n", .{ leak, if (leak != 1) "s" else "" }); } - printer.fmt("\n", .{}); - try slowest.display(printer); - printer.fmt("\n", .{}); + Printer.fmt("\n", .{}); + try slowest.display(); + Printer.fmt("\n", .{}); std.posix.exit(if (fail == 0) 0 else 1); } const Printer = struct { - out: std.fs.File.Writer, - - fn init() Printer { - return .{ - .out = std.io.getStdErr().writer(), - }; + fn fmt(comptime format: []const u8, args: anytype) void { + std.debug.print(format, args); } - fn fmt(self: Printer, comptime format: []const u8, args: anytype) void { - std.fmt.format(self.out, format, args) catch unreachable; - } - - fn status(self: Printer, s: Status, comptime format: []const u8, args: anytype) void { - const color = switch (s) { - .pass => "\x1b[32m", - .fail => "\x1b[31m", - .skip => "\x1b[33m", - else => "", - }; - const out = self.out; - out.writeAll(color) catch @panic("writeAll failed?!"); - std.fmt.format(out, format, args) catch @panic("std.fmt.format failed?!"); - self.fmt("\x1b[0m", .{}); + fn status(s: Status, comptime format: []const u8, args: anytype) void { + switch (s) { + .pass => std.debug.print("\x1b[32m", .{}), + .fail => std.debug.print("\x1b[31m", .{}), + .skip => std.debug.print("\x1b[33m", .{}), + else => {}, + } + std.debug.print(format ++ "\x1b[0m", args); } }; @@ -233,13 +219,13 @@ const SlowTracker = struct { return ns; } - fn display(self: *SlowTracker, printer: Printer) !void { + fn display(self: *SlowTracker) !void { var slowest = self.slowest; const count = slowest.count(); - printer.fmt("Slowest {d} test{s}: \n", .{ count, if (count != 1) "s" else "" }); + Printer.fmt("Slowest {d} test{s}: \n", .{ count, if (count != 1) "s" else "" }); while (slowest.removeMinOrNull()) |info| { const ms = @as(f64, @floatFromInt(info.ns)) / 1_000_000.0; - printer.fmt(" {d:.2}ms\t{s}\n", .{ ms, info.name }); + Printer.fmt(" {d:.2}ms\t{s}\n", .{ ms, info.name }); } } -- 2.51.2 From fd6894800a8cfd559466373b4c120b5b3ad44fb7 Mon Sep 17 00:00:00 2001 From: Karl Seguin Date: Thu, 24 Jul 2025 08:01:18 +0800 Subject: [PATCH 2/5] Make testing.init take an option parameter Allow port to be specified. https://github.com/karlseguin/websocket.zig/issues/71 ``` wt.init(.{.port = 3233}); ``` --- readme.md | 2 +- src/testing.zig | 7 +++++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/readme.md b/readme.md index 1c18fd3..b7c4fcb 100644 --- a/readme.md +++ b/readme.md @@ -403,7 +403,7 @@ The library comes with some helpers for testing. const wt = @import("websocket").testing; test "handler: echo" { - var wtt = wt.init(); + var wtt = wt.init(.{}); defer wtt.deinit(); // create an instance of your handler (however you want) diff --git a/src/testing.zig b/src/testing.zig index ed617c6..d37362c 100644 --- a/src/testing.zig +++ b/src/testing.zig @@ -16,7 +16,10 @@ pub const Testing = struct { received: std.ArrayList(ws.Message), received_index: usize, - fn init() Testing { + const Opts = struct { + port: ?u16 = null, + }; + fn init(opts: Opts) Testing { const arena = t.allocator.create(std.heap.ArenaAllocator) catch unreachable; errdefer t.allocator.destroy(arena); @@ -49,7 +52,7 @@ pub const Testing = struct { ._closed = false, .started = 0, .stream = pair.server, - .address = std.net.Address.parseIp("127.0.0.1", 0) catch unreachable, + .address = std.net.Address.parseIp("127.0.0.1", opts.port orelse 0) catch unreachable, }, .reader = reader, .received = std.ArrayList(ws.Message).init(aa), -- 2.51.2 From 069cc04c22fee4e8f77ee7093257cf479ce19e18 Mon Sep 17 00:00:00 2001 From: Karl Seguin Date: Fri, 25 Jul 2025 20:36:59 +0800 Subject: [PATCH 3/5] Pass configured port to underlying socketpair. Make testing.ensureMessage public. https://github.com/karlseguin/websocket.zig/issues/71 --- src/client/client.zig | 20 ++++++++++---------- src/proto.zig | 6 +++--- src/t.zig | 8 ++++++-- src/testing.zig | 7 ++++--- 4 files changed, 23 insertions(+), 18 deletions(-) diff --git a/src/client/client.zig b/src/client/client.zig index ce69b2b..7e502f4 100644 --- a/src/client/client.zig +++ b/src/client/client.zig @@ -578,7 +578,7 @@ const t = @import("../t.zig"); test "Client: handshake" { { // empty response - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("\r\n\r\n"); @@ -589,7 +589,7 @@ test "Client: handshake" { { // invalid websocket response - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 200 OK\r\n\r\n"); @@ -600,7 +600,7 @@ test "Client: handshake" { { // missing upgrade header - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\n\r\n"); @@ -611,7 +611,7 @@ test "Client: handshake" { { // wrong upgrade header - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: nope\r\n\r\n"); @@ -622,7 +622,7 @@ test "Client: handshake" { { // missing connection header - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\n\r\n"); @@ -633,7 +633,7 @@ test "Client: handshake" { { // wrong connection header - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: something\r\n\r\n"); @@ -644,7 +644,7 @@ test "Client: handshake" { { // missing Sec-Websocket-Accept header - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\nConnection: upgrade\r\n\r\n"); @@ -655,7 +655,7 @@ test "Client: handshake" { { // wrong Sec-Websocket-Accept header - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: hack\r\n\r\n"); @@ -666,7 +666,7 @@ test "Client: handshake" { { // ok for successful - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: C/0nmHhBztSRGR1CwL6Tf4ZjwpY=\r\n\r\n"); @@ -678,7 +678,7 @@ test "Client: handshake" { { // ok for successful, with overread - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.client.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: C/0nmHhBztSRGR1CwL6Tf4ZjwpY=\r\n\r\nSome Random Data Which is Part Of the Next Message"); diff --git a/src/proto.zig b/src/proto.zig index a3c3eca..e981f08 100644 --- a/src/proto.zig +++ b/src/proto.zig @@ -620,7 +620,7 @@ test "mask" { test "Reader: read too large" { defer t.reset(); - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); pair.textFrame(true, "hello world"); pair.sendBuf(); @@ -633,7 +633,7 @@ test "Reader: read too large" { test "Reader: read too large over multiple fragments" { defer t.reset(); - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); pair.textFrame(false, "hello world"); pair.cont(false, " !!!_!!! "); @@ -648,7 +648,7 @@ test "Reader: read too large over multiple fragments" { test "Reader: exact read into static with no overflow" { defer t.reset(); - var pair = t.SocketPair.init(); + var pair = t.SocketPair.init(.{}); defer pair.deinit(); pair.textFrame(true, "hello!"); pair.sendBuf(); diff --git a/src/t.zig b/src/t.zig index fe34e79..61e0a8f 100644 --- a/src/t.zig +++ b/src/t.zig @@ -148,8 +148,12 @@ pub const SocketPair = struct { client: std.net.Stream, server: std.net.Stream, - pub fn init() SocketPair { - var address = std.net.Address.parseIp("127.0.0.1", 0) catch unreachable; + const Opts = struct { + port: ?u16 = null, + }; + + pub fn init(opts: Opts) SocketPair { + var address = std.net.Address.parseIp("127.0.0.1", opts.port orelse 0) catch unreachable; var address_len = address.getOsSockLen(); const listener = posix.socket(address.any.family, posix.SOCK.STREAM | posix.SOCK.CLOEXEC, posix.IPPROTO.TCP) catch unreachable; diff --git a/src/testing.zig b/src/testing.zig index d37362c..f9c279f 100644 --- a/src/testing.zig +++ b/src/testing.zig @@ -26,7 +26,8 @@ pub const Testing = struct { arena.* = std.heap.ArenaAllocator.init(t.allocator); errdefer arena.deinit(); - const pair = t.SocketPair.init(); + const port = opts.port orelse 0; + const pair = t.SocketPair.init(.{.port = port}); const timeout = std.mem.toBytes(std.posix.timeval{ .sec = 0, .usec = 50_000, @@ -52,7 +53,7 @@ pub const Testing = struct { ._closed = false, .started = 0, .stream = pair.server, - .address = std.net.Address.parseIp("127.0.0.1", opts.port orelse 0) catch unreachable, + .address = std.net.Address.parseIp("127.0.0.1", port) catch unreachable, }, .reader = reader, .received = std.ArrayList(ws.Message).init(aa), @@ -97,7 +98,7 @@ pub const Testing = struct { // we have a 50ms timeout on this socket. It's all localhost. We expect // to be able to read messages in that time. - fn ensureMessage(self: *Testing) !void { + pub fn ensureMessage(self: *Testing) !void { if (self.received_index < self.received.items.len) { return; } -- 2.51.2 From 0960332e63793eef3815279165538b8b7458bc95 Mon Sep 17 00:00:00 2001 From: Karl Seguin Date: Sat, 26 Jul 2025 15:40:20 +0800 Subject: [PATCH 4/5] createReply allow null headers --- src/server/handshake.zig | 183 ++++++++++++++++++++------------------- 1 file changed, 92 insertions(+), 91 deletions(-) diff --git a/src/server/handshake.zig b/src/server/handshake.zig index 76d401b..98b4e48 100644 --- a/src/server/handshake.zig +++ b/src/server/handshake.zig @@ -122,7 +122,7 @@ pub const Handshake = struct { }; } - pub fn createReply(key: []const u8, headers: *const KeyValue, compression: ?websocket.Compression, buf: []u8) ![]const u8 { + pub fn createReply(key: []const u8, headers_: ?*KeyValue, compression: ?websocket.Compression, buf: []u8) ![]const u8 { const HEADER = "HTTP/1.1 101 Switching Protocols\r\n" ++ "Upgrade: websocket\r\n" ++ @@ -166,8 +166,10 @@ pub const Handshake = struct { } } - for (headers.keys[0..headers.len], headers.values[0..headers.len]) |k, v| { - pos += (try std.fmt.bufPrint(buf[pos..], "\r\n{s}: {s}", .{ k, v })).len; + if (headers_) |headers| { + for (headers.keys[0..headers.len], headers.values[0..headers.len]) |k, v| { + pos += (try std.fmt.bufPrint(buf[pos..], "\r\n{s}: {s}", .{ k, v })).len; + } } const end = pos + 4; @@ -242,10 +244,10 @@ pub const Handshake = struct { const buf = try allocator.alloc(u8, pool.buffer_size); errdefer allocator.free(buf); - const req_headers = try KeyValue.init(allocator, pool.max_req_headers); + const req_headers = try Handshake.KeyValue.init(allocator, pool.max_req_headers); errdefer req_headers.deinit(allocator); - const res_headers = try KeyValue.init(allocator, pool.max_res_headers); + const res_headers = try Handshake.KeyValue.init(allocator, pool.max_res_headers); errdefer res_headers.deinit(allocator); return .{ @@ -270,96 +272,96 @@ pub const Handshake = struct { self.pool.release(self); } }; -}; -pub const KeyValue = struct { - len: usize, - keys: [][]const u8, - values: [][]const u8, + pub const KeyValue = struct { + len: usize, + keys: [][]const u8, + values: [][]const u8, - fn init(allocator: Allocator, max: usize) !KeyValue { - const keys = try allocator.alloc([]const u8, max); - errdefer allocator.free(keys); + fn init(allocator: Allocator, max: usize) !KeyValue { + const keys = try allocator.alloc([]const u8, max); + errdefer allocator.free(keys); - const values = try allocator.alloc([]const u8, max); - errdefer allocator.free(values); + const values = try allocator.alloc([]const u8, max); + errdefer allocator.free(values); - return .{ - .len = 0, - .keys = keys, - .values = values, - }; - } - - fn deinit(self: *const KeyValue, allocator: Allocator) void { - allocator.free(self.keys); - allocator.free(self.values); - } - - pub fn add(self: *KeyValue, key: []const u8, value: []const u8) void { - const len = self.len; - var keys = self.keys; - if (len == keys.len) { - return; + return .{ + .len = 0, + .keys = keys, + .values = values, + }; } - keys[len] = key; - self.values[len] = value; - self.len = len + 1; - } + fn deinit(self: *const KeyValue, allocator: Allocator) void { + allocator.free(self.keys); + allocator.free(self.values); + } - pub fn get(self: *const KeyValue, needle: []const u8) ?[]const u8 { - const keys = self.keys[0..self.len]; - loop: for (keys, 0..) |key, i| { - // This is largely a reminder to myself that std.mem.eql isn't - // particularly fast. Here we at least avoid the 1 extra ptr - // equality check that std.mem.eql does, but we could do better - // TODO: monitor https://github.com/ziglang/zig/issues/8689 - if (needle.len != key.len) { - continue; + pub fn add(self: *KeyValue, key: []const u8, value: []const u8) void { + const len = self.len; + var keys = self.keys; + if (len == keys.len) { + return; } - for (needle, key) |n, k| { - if (n != k) { - continue :loop; + + keys[len] = key; + self.values[len] = value; + self.len = len + 1; + } + + pub fn get(self: *const KeyValue, needle: []const u8) ?[]const u8 { + const keys = self.keys[0..self.len]; + loop: for (keys, 0..) |key, i| { + // This is largely a reminder to myself that std.mem.eql isn't + // particularly fast. Here we at least avoid the 1 extra ptr + // equality check that std.mem.eql does, but we could do better + // TODO: monitor https://github.com/ziglang/zig/issues/8689 + if (needle.len != key.len) { + continue; } + for (needle, key) |n, k| { + if (n != k) { + continue :loop; + } + } + return self.values[i]; } - return self.values[i]; + + return null; } - return null; - } + pub fn iterator(self: *const KeyValue) Iterator { + const len = self.len; + return .{ + .pos = 0, + .keys = self.keys[0..len], + .values = self.values[0..len], + }; + } - pub fn iterator(self: *const KeyValue) Iterator { - const len = self.len; - return .{ - .pos = 0, - .keys = self.keys[0..len], - .values = self.values[0..len], - }; - } + pub const Iterator = struct { + pos: usize, + keys: [][]const u8, + values: [][]const u8, - pub const Iterator = struct { - pos: usize, - keys: [][]const u8, - values: [][]const u8, + const KV = struct { + key: []const u8, + value: []const u8, + }; - const KV = struct { - key: []const u8, - value: []const u8, - }; + pub fn next(self: *Iterator) ?KV { + const pos = self.pos; + if (pos == self.keys.len) { + return null; + } - pub fn next(self: *Iterator) ?KV { - const pos = self.pos; - if (pos == self.keys.len) { - return null; + self.pos = pos + 1; + return .{ + .key = self.keys[pos], + .value = self.values[pos], + }; } - - self.pos = pos + 1; - return .{ - .key = self.keys[pos], - .value = self.values[pos], - }; - } + }; }; }; @@ -525,8 +527,6 @@ test "handshake: parse" { test "handshake: reply" { var buf: [512]u8 = undefined; - var res_headers = try KeyValue.init(t.allocator, 2); - defer res_headers.deinit(t.allocator); { // no compression @@ -535,7 +535,7 @@ test "handshake: reply" { "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, null, &buf)); + try t.expectString(expected, try Handshake.createReply("this is my key", null, null, &buf)); } { @@ -546,7 +546,7 @@ test "handshake: reply" { "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ "Sec-WebSocket-Extensions: permessage-deflate\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{}, &buf)); + try t.expectString(expected, try Handshake.createReply("this is my key", null, .{}, &buf)); } { @@ -557,14 +557,17 @@ test "handshake: reply" { "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{ + try t.expectString(expected, try Handshake.createReply("this is my key", null, .{ .client_no_context_takeover = true, .server_no_context_takeover = true, }, &buf)); } // With custom headers + var res_headers = try Handshake.KeyValue.init(t.allocator, 2); + defer res_headers.deinit(t.allocator); res_headers.add("Set-Cookie", "Yummy!"); + { // no compression const expected = @@ -583,9 +586,8 @@ test "handshake: reply" { "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Sec-WebSocket-Extensions: permessage-deflate\r\n" ++ - "Set-Cookie: Yummy!\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{}, &buf)); + "Sec-WebSocket-Extensions: permessage-deflate\r\n\r\n"; + try t.expectString(expected, try Handshake.createReply("this is my key", null, .{}, &buf)); } { @@ -595,9 +597,8 @@ test "handshake: reply" { "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n" ++ - "Set-Cookie: Yummy!\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{ + "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n\r\n"; + try t.expectString(expected, try Handshake.createReply("this is my key", null, .{ .client_no_context_takeover = true, .server_no_context_takeover = true, }, &buf)); @@ -606,7 +607,7 @@ test "handshake: reply" { test "KeyValue: get" { const allocator = t.allocator; - var kv = try KeyValue.init(allocator, 2); + var kv = try Handshake.KeyValue.init(allocator, 2); defer kv.deinit(t.allocator); var key = "content-type".*; @@ -621,7 +622,7 @@ test "KeyValue: get" { } test "KeyValue: ignores beyond max" { - var kv = try KeyValue.init(t.allocator, 2); + var kv = try Handshake.KeyValue.init(t.allocator, 2); defer kv.deinit(t.allocator); var n1 = "content-length".*; -- 2.51.2 From abbdfc0e7f391aae19d3f162d96c6280eb69c305 Mon Sep 17 00:00:00 2001 From: Karl Seguin Date: Sat, 23 Aug 2025 19:35:39 +0800 Subject: [PATCH 5/5] zig 0.15 support (temp removal of compression) --- readme.md | 6 - src/buffer.zig | 23 ++++ src/client/client.zig | 101 +++++++++++++---- src/proto.zig | 98 ++++++---------- src/server/handshake.zig | 75 +++---------- src/server/server.zig | 168 +++++++++++++--------------- src/t.zig | 8 +- src/testing.zig | 2 +- src/websocket.zig | 6 +- support/autobahn/client/build.zig | 8 +- support/autobahn/server/build.zig | 8 +- support/autobahn/server/config.json | 5 +- support/autobahn/server/main.zig | 6 +- 13 files changed, 250 insertions(+), 264 deletions(-) diff --git a/readme.md b/readme.md index b7c4fcb..e35ee9d 100644 --- a/readme.md +++ b/readme.md @@ -340,12 +340,6 @@ pub const Config = struct { // is freed after each message. // true = more memory, but fewer allocations retain_write_buffer: bool = true, - - // Advanced options that are part of the permessage-deflate specification. - // You can set these to true to try and save a bit of memory. But if you - // want to save memory, don't use compression at all. - client_no_context_takeover: bool = false, - server_no_context_takeover: bool = false, }; } ``` diff --git a/src/buffer.zig b/src/buffer.zig index e9ed5d8..e1cb78d 100644 --- a/src/buffer.zig +++ b/src/buffer.zig @@ -18,6 +18,21 @@ pub const Writer = struct { pos: usize = 0, pooled: bool, provider: *Provider, + interface: std.Io.Writer, + + pub fn init(buf: []u8, pooled: bool, provider: *Provider, dumb: []u8) Writer { + return .{ + .buf = buf, + .pooled = pooled, + .provider = provider, + .interface = .{ + .buffer = dumb, + .vtable = &.{ + .drain = drain, + }, + }, + }; + } pub fn deinit(self: *Writer) void { if (self.pooled) { @@ -27,6 +42,14 @@ pub const Writer = struct { } } + pub fn drain(io_w: *std.io.Writer, data: []const []const u8, splat: usize) error{WriteFailed}!usize { + std.debug.print("drain: {d}\n", .{data[0].len}); + _ = splat; + const self: *Writer = @alignCast(@fieldParentPtr("interface", io_w)); + self.writeAll(data[0]) catch return error.WriteFailed; + return data[0].len; + } + pub fn writeAll(self: *Writer, data: []const u8) !void { const pos = self.pos; const total_len = pos + data.len; diff --git a/src/client/client.zig b/src/client/client.zig index 7e502f4..f2bec5b 100644 --- a/src/client/client.zig +++ b/src/client/client.zig @@ -66,22 +66,9 @@ pub const Client = struct { pub fn init(allocator: Allocator, config: Config) !Client { const net_stream = try net.tcpConnectToHost(allocator, config.host, config.port); - var tls_client: ?tls.Client = null; + var tls_client: ?*TLSClient = null; if (config.tls) { - var own_bundle = false; - var bundle = config.ca_bundle orelse blk: { - own_bundle = true; - var b = Bundle{}; - try b.rescan(allocator); - break :blk b; - }; - defer if (own_bundle) { - bundle.deinit(allocator); - }; - tls_client = try tls.Client.init(net_stream, .{ - .host = .{ .explicit = config.host }, - .ca = .{ .bundle = bundle }, - }); + tls_client = try TLSClient.init(allocator, net_stream, &config); } const stream = Stream.init(net_stream, tls_client); @@ -340,9 +327,9 @@ pub const Client = struct { // wraps a net.Stream and optional a tls.Client pub const Stream = struct { stream: net.Stream, - tls_client: ?tls.Client = null, + tls_client: ?*TLSClient = null, - pub fn init(stream: net.Stream, tls_client: ?tls.Client) Stream { + pub fn init(stream: net.Stream, tls_client: ?*TLSClient) Stream { return .{ .stream = stream, .tls_client = tls_client, @@ -350,8 +337,8 @@ pub const Stream = struct { } pub fn close(self: *Stream) void { - if (self.tls_client) |*tls_client| { - _ = tls_client.writeEnd(self.stream, "", true) catch {}; + if (self.tls_client) |tls_client| { + tls_client.deinit(); } // std.posix.close panics on EBADF @@ -374,15 +361,26 @@ pub const Stream = struct { } pub fn read(self: *Stream, buf: []u8) !usize { - if (self.tls_client) |*tls_client| { - return tls_client.read(self.stream, buf); + if (self.tls_client) |tls_client| { + var w: std.Io.Writer = .fixed(buf); + while (true) { + const n = try tls_client.client.reader.stream(&w, .limited(buf.len)); + if (n != 0) { + return n; + } + } } return self.stream.read(buf); } pub fn writeAll(self: *Stream, data: []const u8) !void { - if (self.tls_client) |*tls_client| { - return tls_client.writeAll(self.stream, data); + if (self.tls_client) |tls_client| { + try tls_client.client.writer.writeAll(data); + // I know this looks silly, but as far as I can tell, this is what + // we need to do. + try tls_client.client.writer.flush(); + try tls_client.stream_writer.interface.flush(); + return; } return self.stream.writeAll(data); } @@ -413,6 +411,63 @@ pub const Stream = struct { } }; +const TLSClient = struct { + client: tls.Client, + stream: net.Stream, + stream_writer: net.Stream.Writer, + stream_reader: net.Stream.Reader, + arena: std.heap.ArenaAllocator, + + fn init(allocator: Allocator, stream: net.Stream, config: *const Client.Config) !*TLSClient { + var arena = std.heap.ArenaAllocator.init(allocator); + errdefer arena.deinit(); + + const aa = arena.allocator(); + + const bundle = config.ca_bundle orelse blk: { + var b = Bundle{}; + try b.rescan(aa); + break :blk b; + }; + + // The TLS input and output have to be max_ciphertext_record_len each. + // It isn't clear to me how big the un-encrypted reader and writer + // need to be. I would think 0, but that will fail an assertion. I + // don't think that it's right that we need 4 buffers, but apparently + // we do. Until i figure this out, using 4 x max_ciphertext_record_len + // seems like the only safe choice. + const buf_len = std.crypto.tls.max_ciphertext_record_len; + var buf = try aa.alloc(u8, buf_len * 4); + + const self = try aa.create(TLSClient); + self.* = .{ + .stream = stream, + .arena = arena, + .client = undefined, + .stream_writer = stream.writer(buf.ptr[0..buf_len][0..buf_len]), + .stream_reader = stream.reader(buf.ptr[buf_len .. 2 * buf_len][0..buf_len]), + }; + + self.client = try tls.Client.init( + self.stream_reader.interface(), + &self.stream_writer.interface, + .{ + .ca = .{ .bundle = bundle }, + .host = .{ .explicit = config.host }, + .read_buffer = buf.ptr[2 * buf_len .. 3 * buf_len][0..buf_len], + .write_buffer = buf.ptr[3 * buf_len .. 4 * buf_len][0..buf_len], + }, + ); + + return self; + } + + fn deinit(self: *TLSClient) void { + _ = self.client.end() catch {}; + self.arena.deinit(); + } +}; + fn generateKey() [16]u8 { if (comptime @import("builtin").is_test) { return [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; diff --git a/src/proto.zig b/src/proto.zig index e981f08..7eeff8d 100644 --- a/src/proto.zig +++ b/src/proto.zig @@ -87,26 +87,15 @@ pub const Reader = struct { // fragment), the state of the fragmented message is maintained here.) fragment: ?Fragmented, - decompressor: ?DecompressorType, + allow_compressed: bool, // if we returned a decompressed message, it's stored here so that we can // cleanup when the user is done with the message decompress_writer: ?buffer.Writer, - // when client_no_context_takeover, we reset the decompressor after every - // message - decompressor_reset: bool, - - const DecompressorType = std.compress.flate.Decompressor(std.io.FixedBufferStream([]const u8).Reader); + const DecompressorType = std.compress.flate.Decompress; pub fn init(static: []u8, large_buffer_provider: *buffer.Provider, compression: ?Compression) Reader { - var decompressor_reset = false; - var decompressor: ?DecompressorType = null; - if (compression) |c| { - decompressor = .{}; - decompressor_reset = c.client_no_context_takeover; - } - return .{ .pos = 0, .start = 0, @@ -115,14 +104,13 @@ pub const Reader = struct { .message_len = 0, .fragment = null, .decompress_writer = null, - .decompressor = decompressor, - .decompressor_reset = decompressor_reset, + .allow_compressed = compression != null, .large_buffer_provider = large_buffer_provider, }; } pub fn deinit(self: *Reader) void { - if (self.fragment) |f| { + if (self.fragment) |*f| { f.deinit(); } if (self.decompress_writer) |*dw| { @@ -265,14 +253,14 @@ pub const Reader = struct { // FIN, RSV1, RSV2, RSV3, OP,OP,OP,OP // RSV2 and RSV3 should never be set, and RSV1 should not be set // when compression is disabled - const rsv_bits: u8 = if (self.decompressor == null or is_continuation) 112 else 48; + const rsv_bits: u8 = if (self.allow_compressed == false or is_continuation) 112 else 48; if (byte1 & rsv_bits != 0) { return error.ReservedFlags; } const compressed = byte1 & 64 == 64; if (compressed) { - if (self.decompressor == null) { + if (self.allow_compressed == false) { return error.CompressionDisabled; } } @@ -371,17 +359,13 @@ pub const Reader = struct { // call to "read" to do this, because we don't know when that'll be. pub fn done(self: *Reader, message_type: Message.Type) void { if (message_type == .text or message_type == .binary) { - if (self.fragment) |f| { + if (self.fragment) |*f| { f.deinit(); self.fragment = null; } if (self.decompress_writer) |*dw| { dw.deinit(); self.decompress_writer = null; - - if (self.decompressor_reset) { - self.decompressor = .{}; - } } } @@ -420,44 +404,26 @@ pub const Reader = struct { fn decompress(self: *Reader, compressed: []const u8) ![]u8 { const provider = self.large_buffer_provider; + var dumb: [32]u8 = undefined; var writer: buffer.Writer = undefined; if (compressed.len < provider.pool_buffer_size) { - writer = .{ - .pooled = true, - .provider = provider, - .buf = try provider.pool.acquireOrCreate(), - }; + const buf = try provider.pool.acquireOrCreate(); + writer = .init(buf, true, provider, &dumb); } else { - writer = .{ - .pooled = false, - .provider = provider, - .buf = try provider.allocator.alloc(u8, @intFromFloat(@as(f64, @floatFromInt(compressed.len)) * 1.25)), - }; + const buf = try provider.allocator.alloc(u8, @intFromFloat(@as(f64, @floatFromInt(compressed.len)) * 1.25)); + writer = .init(buf, false, provider, &dumb); } errdefer writer.deinit(); - var decompressor = &self.decompressor.?; - { - var reader = std.io.fixedBufferStream(compressed); - decompressor.setReader(reader.reader()); - decompressor.decompress(&writer) catch |err| switch (err) { - error.EndOfStream => {}, - else => return error.CompressionError, - }; - } - - { - var reader = std.io.fixedBufferStream(&[_]u8{ 0x00, 0x00, 0xff, 0xff }); - decompressor.setReader(reader.reader()); - decompressor.decompress(&writer) catch |err| switch (err) { - error.EndOfStream => {}, - else => return error.CompressionError, - }; - } + var reader = std.Io.Reader.fixed(compressed); + var decompressor = std.compress.flate.Decompress.init(&reader, .raw, &.{}); + const n = decompressor.reader.streamRemaining(&writer.interface) catch { + return error.CompressionError; + }; self.decompress_writer = writer; - return writer.buf[0..writer.pos]; + return writer.buf[0..n]; } inline fn usingLargeBuffer(self: *const Reader) bool { @@ -480,10 +446,11 @@ const Fragmented = struct { compressed: bool, type: Message.Type, buf: std.ArrayList(u8), + allocator: std.mem.Allocator, pub fn init(bp: *buffer.Provider, compressed: bool, message_type: Message.Type, value: []const u8) !Fragmented { - var buf = std.ArrayList(u8).init(bp.allocator); - try buf.ensureTotalCapacity(value.len * 2); + var buf: std.ArrayList(u8) = .empty; + try buf.ensureTotalCapacity(bp.allocator, value.len * 2); buf.appendSliceAssumeCapacity(value); return .{ @@ -491,18 +458,19 @@ const Fragmented = struct { .type = message_type, .compressed = compressed, .max = bp.max_buffer_size, + .allocator = bp.allocator, }; } - pub fn deinit(self: Fragmented) void { - self.buf.deinit(); + pub fn deinit(self: *Fragmented) void { + self.buf.deinit(self.allocator); } pub fn add(self: *Fragmented, value: []const u8) !void { if (self.buf.items.len + value.len > self.max) { return error.TooLarge; } - try self.buf.appendSlice(value); + try self.buf.appendSlice(self.allocator, value); } // Optimization so that we don't over-allocate on our last frame. @@ -511,7 +479,7 @@ const Fragmented = struct { if (total_len > self.max) { return error.TooLarge; } - try self.buf.ensureTotalCapacityPrecise(total_len); + try self.buf.ensureTotalCapacityPrecise(self.allocator, total_len); self.buf.appendSliceAssumeCapacity(value); return self.buf.items; } @@ -680,7 +648,7 @@ test "Reader: fuzz" { var is_fragmented = false; var fragment_count: usize = 0; - var fragment = std.ArrayList(u8).init(arena); + var fragment: std.ArrayList(u8) = .empty; var i: usize = 0; while (i < MESSAGE_TO_SEND) { @@ -710,7 +678,7 @@ test "Reader: fuzz" { } fragment_count += 1; - try fragment.appendSlice(try arena.dupe(u8, buf)); + try fragment.appendSlice(arena, try arena.dupe(u8, buf)); if (is_fin) { // this was the last message in our fragment @@ -818,20 +786,20 @@ test "Fragmented" { var f = try Fragmented.init(&bp, false, .binary, payload); defer f.deinit(); - var expected = std.ArrayList(u8).init(t.allocator); - defer expected.deinit(); - try expected.appendSlice(payload); + var expected: std.ArrayList(u8) = .empty; + defer expected.deinit(t.allocator); + try expected.appendSlice(t.allocator, payload); const number_of_adds = random.uintAtMost(usize, 30); for (0..number_of_adds) |_| { payload = buf[0 .. random.uintAtMost(usize, 99) + 1]; random.bytes(payload); try f.add(payload); - try expected.appendSlice(payload); + try expected.appendSlice(t.allocator, payload); } payload = buf[0 .. random.uintAtMost(usize, 99) + 1]; random.bytes(payload); - try expected.appendSlice(payload); + try expected.appendSlice(t.allocator, payload); try t.expectString(expected.items, try f.last(payload)); } diff --git a/src/server/handshake.zig b/src/server/handshake.zig index 98b4e48..e5b9f9b 100644 --- a/src/server/handshake.zig +++ b/src/server/handshake.zig @@ -122,7 +122,7 @@ pub const Handshake = struct { }; } - pub fn createReply(key: []const u8, headers_: ?*KeyValue, compression: ?websocket.Compression, buf: []u8) ![]const u8 { + pub fn createReply(key: []const u8, headers_: ?*KeyValue, compression: bool, buf: []u8) ![]const u8 { const HEADER = "HTTP/1.1 101 Switching Protocols\r\n" ++ "Upgrade: websocket\r\n" ++ @@ -145,25 +145,15 @@ pub const Handshake = struct { pos = end; } - if (compression) |c| { - { - const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate"; - const end = pos + permessage_deflate.len; - @memcpy(buf[pos..end], permessage_deflate); - pos = end; - } - if (c.server_no_context_takeover) { - const server_no_context_takeover = "; server_no_context_takeover"; - const end = pos + server_no_context_takeover.len; - @memcpy(buf[pos..end], server_no_context_takeover); - pos = end; - } - if (c.client_no_context_takeover) { - const client_no_context_takeover = "; client_no_context_takeover"; - const end = pos + client_no_context_takeover.len; - @memcpy(buf[pos..end], client_no_context_takeover); - pos = end; - } + if (compression) { + const permessage_deflate = + "\r\nSec-WebSocket-Extensions: permessage-deflate" ++ + "; server_no_context_takeover" ++ + "; client_no_context_takeover"; + + const end = pos + permessage_deflate.len; + @memcpy(buf[pos..end], permessage_deflate); + pos = end; } if (headers_) |headers| { @@ -535,18 +525,7 @@ test "handshake: reply" { "Upgrade: websocket\r\n" ++ "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", null, null, &buf)); - } - - { - // compression - const expected = - "HTTP/1.1 101 Switching Protocols\r\n" ++ - "Upgrade: websocket\r\n" ++ - "Connection: upgrade\r\n" ++ - "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Sec-WebSocket-Extensions: permessage-deflate\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", null, .{}, &buf)); + try t.expectString(expected, try Handshake.createReply("this is my key", null, false, &buf)); } { @@ -557,10 +536,7 @@ test "handshake: reply" { "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", null, .{ - .client_no_context_takeover = true, - .server_no_context_takeover = true, - }, &buf)); + try t.expectString(expected, try Handshake.createReply("this is my key", null, true, &buf)); } // With custom headers @@ -576,32 +552,7 @@ test "handshake: reply" { "Connection: upgrade\r\n" ++ "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ "Set-Cookie: Yummy!\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, null, &buf)); - } - - { - // compression - const expected = - "HTTP/1.1 101 Switching Protocols\r\n" ++ - "Upgrade: websocket\r\n" ++ - "Connection: upgrade\r\n" ++ - "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Sec-WebSocket-Extensions: permessage-deflate\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", null, .{}, &buf)); - } - - { - // compression - const expected = - "HTTP/1.1 101 Switching Protocols\r\n" ++ - "Upgrade: websocket\r\n" ++ - "Connection: upgrade\r\n" ++ - "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ - "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n\r\n"; - try t.expectString(expected, try Handshake.createReply("this is my key", null, .{ - .client_no_context_takeover = true, - .server_no_context_takeover = true, - }, &buf)); + try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, false, &buf)); } } diff --git a/src/server/server.zig b/src/server/server.zig index 1c33e6e..a16447f 100644 --- a/src/server/server.zig +++ b/src/server/server.zig @@ -110,6 +110,12 @@ pub fn Server(comptime H: type) type { } } + var c = config; + if (c.compression != null) { + log.warn("Compression is disabled as part of the 0.15 upgrade. I do hope to re-enable it soon.", .{}); + c.compression = null; + } + const signals = try allocator.alloc(posix.fd_t, config.workerCount()); errdefer allocator.free(signals); @@ -121,7 +127,7 @@ pub fn Server(comptime H: type) type { ._cond = .{}, ._state = state, ._signals = signals, - .config = config, + .config = c, .allocator = allocator, }; } @@ -253,7 +259,7 @@ pub fn Server(comptime H: type) type { started += 1; } - log.info("starting nonblocking worker to listen on {}", .{address}); + log.info("starting nonblocking worker to listen on {f}", .{address}); // in case startInNewThread is waiting self._cond.signal(); @@ -344,11 +350,11 @@ pub fn Blocking(comptime H: type) type { log.err("failed to accept socket: {}", .{err}); continue; }; - log.debug("({}) connected", .{address}); + log.debug("({f}) connected", .{address}); const thread = std.Thread.spawn(.{}, Self.handleConnection, .{ self, socket, address, ctx }) catch |err| { posix.close(socket); - log.err("({}) failed to spawn connection thread: {}", .{ address, err }); + log.err("({f}) failed to spawn connection thread: {}", .{ address, err }); continue; }; thread.detach(); @@ -359,7 +365,7 @@ pub fn Blocking(comptime H: type) type { // Wrapper around _handleConnection so that we can handle erros fn handleConnection(self: *Self, socket: posix.socket_t, address: net.Address, ctx: anytype) void { self._handleConnection(socket, address, ctx) catch |err| { - log.err("({}) uncaught error in connection handler: {}", .{ address, err }); + log.err("({f}) uncaught error in connection handler: {}", .{ address, err }); }; } @@ -382,7 +388,9 @@ pub fn Blocking(comptime H: type) type { } if (hc.handler != null) { // if we have a handler, the our handshake completed - try conn_manager.setupCompression(hc, compression); + if (compression) { + try conn_manager.setupCompression(hc); + } break; } if (timestamp() > deadline) { @@ -581,7 +589,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { // this connection has timed out. Don't use self.cleanup since there's // a bunch of stuff we can assume here..like there's no handler or reader conn.closeSocket(); - log.debug("({}) handshake timeout", .{conn.address}); + log.debug("({f}) handshake timeout", .{conn.address}); if (hc.handshake) |h| { h.release(); } @@ -606,7 +614,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { return if (err == error.WouldBlock) {} else err; }; - log.debug("({}) connected", .{address}); + log.debug("({f}) connected", .{address}); { errdefer posix.close(socket); @@ -642,7 +650,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { var success = false; if (hc.handler == null) { success = self.dataForHandshake(hc) catch |err| blk: { - log.err("({any}) error processing handshake: {}", .{ hc.conn.address, err }); + log.err("({f}) error processing handshake: {}", .{ hc.conn.address, err }); break :blk false; }; } else { @@ -662,7 +670,7 @@ fn NonBlocking(comptime H: type, comptime C: type) type { self.base.cleanupConn(hc); } else { self.loop.monitorRead(hc, true) catch |err| { - log.debug("({}) failed to add read event monitor: {}", .{ conn.address, err }); + log.debug("({f}) failed to add read event monitor: {}", .{ conn.address, err }); conn.closeSocket(); self.base.cleanupConn(hc); }; @@ -681,7 +689,9 @@ fn NonBlocking(comptime H: type, comptime C: type) type { conn_manager.inactive(hc); } - try conn_manager.setupCompression(hc, compression); + if (compression) { + try conn_manager.setupCompression(hc); + } return true; } }; @@ -740,7 +750,7 @@ fn NonBlockingBase(comptime H: type, comptime MANAGE_HS: bool) type { pub fn dataAvailable(self: *Self, hc: *HandlerConn(H), thread_buf: []u8) bool { return self._dataAvailable(hc, thread_buf) catch |err| { - log.err("({any}) error processing client message: {}", .{ hc.conn.address, err }); + log.err("({f}) error processing client message: {}", .{ hc.conn.address, err }); return false; }; } @@ -1002,8 +1012,11 @@ pub fn Worker(comptime H: type) type { return self.worker.conn_manager.compression != null; } - pub fn setupConnection(self: *Self, hc: *HandlerConn(H), agreed: ?Compression) !void { - return self.worker.conn_manager.setupCompression(hc, agreed); + pub fn setupConnection( + self: *Self, + hc: *HandlerConn(H), + ) !void { + return self.worker.conn_manager.setupCompression(hc); } pub fn shutdown(self: *Self) void { @@ -1223,37 +1236,29 @@ pub fn ConnManager(comptime H: type, comptime MANAGE_HS: bool) type { self.lock.unlock(); } - fn setupCompression(self: *Self, hc: *HandlerConn(H), agreed_: ?Compression) !void { - const agreed = agreed_ orelse { - return; - }; + fn setupCompression(self: *Self, hc: *HandlerConn(H)) !void { + const config = self.compression orelse return; - const configured = self.compression.?; - const merged = Compression{ - .write_threshold = configured.write_threshold, - .retain_write_buffer = configured.retain_write_buffer, - .client_no_context_takeover = agreed.client_no_context_takeover, - .server_no_context_takeover = agreed.server_no_context_takeover, - }; - hc.compression = merged; + hc.compression = config; - if (merged.write_threshold == null) { - // and we have a write threshold, we need to setup our - // connection's compression (read compression is configured in our reader) + if (config.write_threshold == null) { + // if write_treshold is null, then we never want to compress + // outgoing messages. We don't need to set the conn.compression + // field. + // We'll still [potentially] decompress incoming messages, but + // that's set on the proto. return; } - var compression = try self.compression_pool.create(); + const compression = try self.compression_pool.create(); errdefer self.compression_pool.destroy(compression); compression.* = .{ - .compressor = undefined, - .write_treshold = merged.write_threshold.?, - .reset = merged.server_no_context_takeover, - .retain_writer = merged.retain_write_buffer, - .writer = std.ArrayList(u8).init(self.allocator), + .allocator = self.allocator, + .write_treshold = config.write_threshold.?, + .retain_writer = config.retain_write_buffer, + .writer = std.Io.Writer.Allocating.init(self.allocator), }; - compression.compressor = try Conn.Compression.Type.init(compression.writer.writer(), .{}); hc.conn.compression = compression; } @@ -1299,13 +1304,10 @@ pub const Conn = struct { compression: ?*Conn.Compression = null, const Compression = struct { - reset: bool, + allocator: Allocator, retain_writer: bool, write_treshold: usize, - compressor: Type, - writer: std.ArrayList(u8), - - const Type = std.compress.flate.Compressor(std.ArrayList(u8).Writer); + writer: std.Io.Writer.Allocating, }; pub fn isClosed(self: *Conn) bool { @@ -1377,24 +1379,21 @@ pub const Conn = struct { if (data.len >= c.write_treshold) { compressed = true; - var writer = &c.writer; - var compressor = &c.compressor; - var fbs = std.io.fixedBufferStream(data); - _ = try compressor.compress(fbs.reader()); - try compressor.flush(); - payload = writer.items[0 .. writer.items.len - 4]; - - if (c.reset) { - c.compressor = try Conn.Compression.Type.init(writer.writer(), .{}); - } + var compressor = std.compress.flate.Compress.init(&c.writer.writer, &.{}, .{}); + try compressor.writer.writeAll(data); + try compressor.writer.flush(); + const all = c.writer.written(); + payload = all[0 .. all.len - 4]; } } + defer if (compressed) { const c = self.compression.?; if (c.retain_writer) { - c.compressor.wrt.context.clearRetainingCapacity(); + c.writer.clearRetainingCapacity(); } else { - c.compressor.wrt.context.clearAndFree(); + c.writer.deinit(); + c.writer = std.Io.Writer.Allocating.init(c.allocator); } }; @@ -1489,14 +1488,14 @@ pub const Conn = struct { }; }; -fn handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) struct { ?Compression, bool } { +fn handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) struct { bool, bool } { return _handleHandshake(H, worker, hc, ctx) catch |err| { - log.warn("({any}) uncaugh error processing handshake: {}", .{ hc.conn.address, err }); - return .{ null, false }; + log.warn("({f}) uncaugh error processing handshake: {}", .{ hc.conn.address, err }); + return .{ false, false }; }; } -fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) !struct { ?Compression, bool } { +fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) !struct { bool, bool } { std.debug.assert(hc.handler == null); var state = hc.handshake orelse blk: { @@ -1510,47 +1509,38 @@ fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: const len = state.len; if (len == buf.len) { - log.warn("({any}) handshake request exceeded maximum configured size ({d})", .{ conn.address, buf.len }); - return .{ null, false }; + log.warn("({f}) handshake request exceeded maximum configured size ({d})", .{ conn.address, buf.len }); + return .{ false, false }; } const n = posix.read(hc.socket, buf[len..]) catch |err| { switch (err) { - error.BrokenPipe, error.ConnectionResetByPeer => log.debug("({any}) handshake connection closed: {}", .{ conn.address, err }), + error.BrokenPipe, error.ConnectionResetByPeer => log.debug("({f}) handshake connection closed: {}", .{ conn.address, err }), error.WouldBlock => { std.debug.assert(blockingMode()); - log.debug("({any}) handshake timeout", .{conn.address}); + log.debug("({f}) handshake timeout", .{conn.address}); }, - else => log.warn("({any}) handshake error reading from socket: {}", .{ conn.address, err }), + else => log.warn("({f}) handshake error reading from socket: {}", .{ conn.address, err }), } - return .{ null, false }; + return .{ false, false }; }; if (n == 0) { - log.debug("({any}) handshake connection closed", .{conn.address}); - return .{ null, false }; + log.debug("({f}) handshake connection closed", .{conn.address}); + return .{ false, false }; } state.len = len + n; var handshake = Handshake.parse(state) catch |err| { - log.debug("({any}) error parsing handshake: {}", .{ conn.address, err }); + log.debug("({f}) error parsing handshake: {}", .{ conn.address, err }); respondToHandshakeError(conn, err); - return .{ null, false }; + return .{ false, false }; } orelse { // we need more data - return .{ null, true }; + return .{ false, true }; }; - var agreed_compression: ?Compression = null; - if (worker.compression) |configured_compression| { - if (handshake.compression) |request_compression| { - agreed_compression = .{ - .client_no_context_takeover = configured_compression.client_no_context_takeover or request_compression.client_no_context_takeover, - .server_no_context_takeover = configured_compression.server_no_context_takeover or request_compression.server_no_context_takeover, - }; - } - } - + const compression = handshake.compression != null and worker.compression != null; defer state.release(); hc.handshake = null; @@ -1563,33 +1553,33 @@ fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: } else { respondToHandshakeError(conn, err); } - log.debug("({}) " ++ @typeName(H) ++ ".init rejected request {}", .{ conn.address, err }); - return .{ null, false }; + log.debug("({f}) " ++ @typeName(H) ++ ".init rejected request {}", .{ conn.address, err }); + return .{ false, false }; }; hc.handler = handler; var reply_buf: [2048]u8 = undefined; - const handshake_reply = try Handshake.createReply(handshake.key, handshake.res_headers, agreed_compression, &reply_buf); + const handshake_reply = try Handshake.createReply(handshake.key, handshake.res_headers, compression, &reply_buf); try conn.writeFramed(handshake_reply); if (comptime std.meta.hasFn(H, "afterInit")) { const params = @typeInfo(@TypeOf(H.afterInit)).@"fn".params; const res = if (params.len == 1) hc.handler.?.afterInit() else hc.handler.?.afterInit(ctx); res catch |err| { - log.debug("({}) " ++ @typeName(H) ++ ".afterInit error: {}", .{ conn.address, err }); + log.debug("({f}) " ++ @typeName(H) ++ ".afterInit error: {}", .{ conn.address, err }); return .{ null, false }; }; } - log.debug("({}) connection successfully upgraded", .{conn.address}); - return .{ agreed_compression, true }; + log.debug("({f}) connection successfully upgraded", .{conn.address}); + return .{ compression, true }; } fn handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator, fba: *FixedBufferAllocator) bool { std.debug.assert(hc.handshake == null); return _handleClientData(H, hc, allocator, fba) catch |err| { - log.warn("({any}) uncaugh error handling incoming data: {}", .{ hc.conn.address, err }); + log.warn("({f}) uncaugh error handling incoming data: {}", .{ hc.conn.address, err }); return false; }; } @@ -1599,8 +1589,8 @@ fn _handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator var reader = &hc.reader.?; reader.fill(conn.stream) catch |err| { switch (err) { - error.BrokenPipe, error.Closed, error.ConnectionResetByPeer => log.debug("({}) connection closed: {}", .{ conn.address, err }), - else => log.warn("({any}) error reading from connection: {}", .{ conn.address, err }), + error.BrokenPipe, error.Closed, error.ConnectionResetByPeer => log.debug("({f}) connection closed: {}", .{ conn.address, err }), + else => log.warn("({f}) error reading from connection: {}", .{ conn.address, err }), } return false; }; @@ -1615,7 +1605,7 @@ fn _handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator error.CompressionError => conn.writeFramed(CLOSE_PROTOCOL_ERROR) catch {}, else => {}, } - log.debug("({any}) invalid websocket packet: {}", .{ conn.address, err }); + log.debug("({f}) invalid websocket packet: {}", .{ conn.address, err }); return false; } orelse { // everything is fine, we just need more data @@ -1625,7 +1615,7 @@ fn _handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator const message_type = message.type; defer reader.done(message_type); - log.debug("({anys}) received {s} message", .{ hc.conn.address, @tagName(message_type) }); + log.debug("({f}) received {s} message", .{ hc.conn.address, @tagName(message_type) }); switch (message_type) { .text, .binary => { const params = @typeInfo(@TypeOf(H.clientMessage)).@"fn".params; diff --git a/src/t.zig b/src/t.zig index 61e0a8f..ebac23c 100644 --- a/src/t.zig +++ b/src/t.zig @@ -35,13 +35,13 @@ pub const Writer = struct { pub fn init() Writer { return .{ .pos = 0, + .buf = .empty, .random = getRandom(), - .buf = std.ArrayList(u8).init(allocator), }; } - pub fn deinit(self: *const Writer) void { - self.buf.deinit(); + pub fn deinit(self: *Writer) void { + self.buf.deinit(allocator); } pub fn ping(self: *Writer) void { @@ -80,7 +80,7 @@ pub const Writer = struct { // 2 byte header + length_of_length + mask + payload_length const needed = 2 + length_of_length + 4 + l; - buf.ensureUnusedCapacity(needed) catch unreachable; + buf.ensureUnusedCapacity(allocator, needed) catch unreachable; if (fin) { buf.appendAssumeCapacity(128 | op_code | reserved); diff --git a/src/testing.zig b/src/testing.zig index f9c279f..9555e8d 100644 --- a/src/testing.zig +++ b/src/testing.zig @@ -27,7 +27,7 @@ pub const Testing = struct { errdefer arena.deinit(); const port = opts.port orelse 0; - const pair = t.SocketPair.init(.{.port = port}); + const pair = t.SocketPair.init(.{ .port = port }); const timeout = std.mem.toBytes(std.posix.timeval{ .sec = 0, .usec = 50_000, diff --git a/src/websocket.zig b/src/websocket.zig index 3a20ec3..fdb822c 100644 --- a/src/websocket.zig +++ b/src/websocket.zig @@ -22,8 +22,10 @@ pub const Handshake = @import("server/handshake.zig").Handshake; pub const Compression = struct { write_threshold: ?usize = null, retain_write_buffer: bool = true, - client_no_context_takeover: bool = false, - server_no_context_takeover: bool = false, + // don't know how to support these with the Zig 0.15 changes. So, for now + // we'll always require these to be true + // client_no_context_takeover: bool = false, + // server_no_context_takeover: bool = false, }; pub fn bufferProvider(allocator: std.mem.Allocator, config: buffer.Config) !buffer.Provider { diff --git a/support/autobahn/client/build.zig b/support/autobahn/client/build.zig index 3b45645..3b6022f 100644 --- a/support/autobahn/client/build.zig +++ b/support/autobahn/client/build.zig @@ -6,9 +6,11 @@ pub fn build(b: *std.Build) void { const exe = b.addExecutable(.{ .name = "autobahn_test_client", - .root_source_file = b.path("main.zig"), - .target = target, - .optimize = optimize, + .root_module = b.createModule(.{ + .root_source_file = b.path("main.zig"), + .target = target, + .optimize = optimize, + }), }); exe.root_module.addImport("websocket", b.dependency("websocket", .{}).module("websocket")); diff --git a/support/autobahn/server/build.zig b/support/autobahn/server/build.zig index e529446..0371cff 100644 --- a/support/autobahn/server/build.zig +++ b/support/autobahn/server/build.zig @@ -6,9 +6,11 @@ pub fn build(b: *std.Build) void { const exe = b.addExecutable(.{ .name = "autobahn_test_server", - .root_source_file = b.path("main.zig"), - .target = target, - .optimize = optimize, + .root_module = b.createModule(.{ + .root_source_file = b.path("main.zig"), + .target = target, + .optimize = optimize, + }), }); const websocket = b.dependency("websocket", .{}).module("websocket"); diff --git a/support/autobahn/server/config.json b/support/autobahn/server/config.json index 079e9f8..233b0dd 100644 --- a/support/autobahn/server/config.json +++ b/support/autobahn/server/config.json @@ -2,10 +2,9 @@ "outdir": "/ab/reports/", "options": {"failByDrop": false}, "servers": [ - {"agent": "non-blocking", "url": "ws://host.docker.internal:9224"}, - {"agent": "non-blocking buffer pool", "url": "ws://host.docker.internal:9225"} + {"agent": "non-blocking", "url": "ws://host.docker.internal:9224"} ], - "cases": ["*"], + "cases": ["12.1.1"], "exclude-cases": [], "exclude-agent-cases": {} } diff --git a/support/autobahn/server/main.zig b/support/autobahn/server/main.zig index 2853db6..7b13cdd 100644 --- a/support/autobahn/server/main.zig +++ b/support/autobahn/server/main.zig @@ -7,7 +7,7 @@ const Handshake = websocket.Handshake; const Allocator = std.mem.Allocator; pub const std_options = std.Options{ .log_scope_levels = &[_]std.log.ScopeLevel{ - .{ .scope = .websocket, .level = .warn }, + .{ .scope = .websocket, .level = .debug }, } }; var nonblocking_server: websocket.Server(Handler) = undefined; @@ -22,7 +22,7 @@ pub fn main() !void { std.posix.sigaction(std.posix.SIG.TERM, &.{ .handler = .{ .handler = shutdown }, - .mask = std.posix.empty_sigset, + .mask = std.posix.sigemptyset(), .flags = 0, }, null); } @@ -107,7 +107,7 @@ const Handler = struct { } }; -fn shutdown(_: c_int) callconv(.C) void { +fn shutdown(_: c_int) callconv(.c) void { nonblocking_server.stop(); nonblocking_bp_server.stop(); } -- 2.51.2