A websocket implementation for zig
0

Configure Feed

Select the types of activity you want to include in your feed.

Merge branch 'dev'

Karl Seguin (Aug 23, 2025, 8:26 PM +0800) 59aaa71e 7240830d

+490 -526
+3 -3
build.zig
··· 5 5 const optimize = b.standardOptimizeOption(.{}); 6 6 7 7 const websocket_module = b.addModule("websocket", .{ 8 + .target = target, 9 + .optimize = optimize, 8 10 .root_source_file = b.path("src/websocket.zig"), 9 11 }); 10 12 ··· 17 19 { 18 20 // run tests 19 21 const tests = b.addTest(.{ 20 - .root_source_file = b.path("src/websocket.zig"), 21 - .target = target, 22 - .optimize = optimize, 22 + .root_module = websocket_module, 23 23 .test_runner = .{ .path = b.path("test_runner.zig"), .mode = .simple }, 24 24 }); 25 25 tests.linkLibC();
+4 -10
readme.md
··· 192 192 `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. 193 193 194 194 ### Writer 195 - 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. 195 + 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. 196 196 197 197 ```zig 198 198 // .text or .binary 199 199 var wb = conn.writeBuffer(allocator, .text); 200 200 defer wb.deinit(); 201 - try std.fmt.format(wb.writer(), "it's over {d}!!!", .{9000}); 202 - try wb.flush(); 201 + try wb.interface.print("it's over {d}!!!", .{9000}); 202 + try wb.send(); 203 203 ``` 204 204 205 205 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 211 212 212 var wb = conn.writeBuffer(allocator, .text); 213 213 try std.fmt.format(wb.writer(), "it's over {d}!!!", .{9000}); 214 - try wb.flush(); 214 + try wb.send(); 215 215 } 216 216 ``` 217 217 ··· 340 340 // is freed after each message. 341 341 // true = more memory, but fewer allocations 342 342 retain_write_buffer: bool = true, 343 - 344 - // Advanced options that are part of the permessage-deflate specification. 345 - // You can set these to true to try and save a bit of memory. But if you 346 - // want to save memory, don't use compression at all. 347 - client_no_context_takeover: bool = false, 348 - server_no_context_takeover: bool = false, 349 343 }; 350 344 } 351 345 ```
+23
src/buffer.zig
··· 18 18 pos: usize = 0, 19 19 pooled: bool, 20 20 provider: *Provider, 21 + interface: std.Io.Writer, 22 + 23 + pub fn init(buf: []u8, pooled: bool, provider: *Provider, dumb: []u8) Writer { 24 + return .{ 25 + .buf = buf, 26 + .pooled = pooled, 27 + .provider = provider, 28 + .interface = .{ 29 + .buffer = dumb, 30 + .vtable = &.{ 31 + .drain = drain, 32 + }, 33 + }, 34 + }; 35 + } 21 36 22 37 pub fn deinit(self: *Writer) void { 23 38 if (self.pooled) { ··· 25 40 } else { 26 41 self.provider.allocator.free(self.buf); 27 42 } 43 + } 44 + 45 + pub fn drain(io_w: *std.io.Writer, data: []const []const u8, splat: usize) error{WriteFailed}!usize { 46 + std.debug.print("drain: {d}\n", .{data[0].len}); 47 + _ = splat; 48 + const self: *Writer = @alignCast(@fieldParentPtr("interface", io_w)); 49 + self.writeAll(data[0]) catch return error.WriteFailed; 50 + return data[0].len; 28 51 } 29 52 30 53 pub fn writeAll(self: *Writer, data: []const u8) !void {
+150 -123
src/client/client.zig
··· 6 6 const net = std.net; 7 7 const posix = std.posix; 8 8 const tls = std.crypto.tls; 9 + const log = std.log.scoped(.websocket); 9 10 10 11 const Reader = proto.Reader; 11 12 const Allocator = std.mem.Allocator; ··· 38 39 _reader: Reader, 39 40 _closed: bool, 40 41 _compression_opts: ?CompressionOpts, 41 - _compression: ?*Client.Compression = null, 42 + _compression: ?Client.Compression = null, 42 43 43 44 // When creating a client, we can either be given a BufferProvider or create 44 45 // one ourselves. If we create it ourselves (in init), we "own" it and must ··· 70 71 }; 71 72 72 73 const Compression = struct { 73 - reset: bool, 74 + allocator: Allocator, 74 75 retain_writer: bool, 75 76 write_treshold: usize, 76 - compressor: Type, 77 - writer: std.ArrayList(u8), 78 - 79 - const Type = std.compress.flate.Compressor(std.ArrayList(u8).Writer); 77 + writer: std.Io.Writer.Allocating, 80 78 }; 81 79 82 80 pub fn init(allocator: Allocator, config: Config) !Client { 81 + if (config.compression != null) { 82 + log.err("Compression is disabled as part of the 0.15 upgrade. I do hope to re-enable it soon.", .{}); 83 + return error.InvalidConfiguraion; 84 + } 85 + 83 86 const net_stream = try net.tcpConnectToHost(allocator, config.host, config.port); 84 87 85 - var tls_client: ?tls.Client = null; 88 + var tls_client: ?*TLSClient = null; 86 89 if (config.tls) { 87 - var own_bundle = false; 88 - var bundle = config.ca_bundle orelse blk: { 89 - own_bundle = true; 90 - var b = Bundle{}; 91 - try b.rescan(allocator); 92 - break :blk b; 93 - }; 94 - defer if (own_bundle) { 95 - bundle.deinit(allocator); 96 - }; 97 - tls_client = try tls.Client.init(net_stream, .{ 98 - .host = .{ .explicit = config.host }, 99 - .ca = .{ .bundle = bundle }, 100 - }); 90 + tls_client = try TLSClient.init(allocator, net_stream, &config); 101 91 } 102 92 const stream = Stream.init(net_stream, tls_client); 103 93 ··· 131 121 return .{ 132 122 .stream = stream, 133 123 ._closed = false, 134 - ._compression_opts = config.compression, 135 124 ._own_bp = own_bp, 136 125 ._mask_fn = config.mask_fn, 126 + ._compression_opts = null, //TODO: ZIG 0.15 137 127 ._reader = Reader.init(reader_buf, buffer_provider, null), 138 128 }; 139 129 } ··· 151 141 larger_buffer_provider.deinit(); 152 142 allocator.destroy(larger_buffer_provider); 153 143 } 154 - 155 - if (self._compression) |compression| { 156 - allocator.destroy(compression); 157 - } 158 144 } 159 145 160 146 pub fn handshake(self: *Client, path: []const u8, opts: HandshakeOpts) !void { ··· 170 156 break :blk std.base64.standard.Encoder.encode(&encoded_key, &bin_key); 171 157 }; 172 158 173 - try sendHandshake(path, key, buf, &opts, self._compression_opts, stream); 159 + try sendHandshake(path, key, buf, &opts, self._compression_opts != null, stream); 174 160 175 - const res = try HandShakeReply.read(buf, key, &opts, self._compression_opts, stream); 161 + const res = try HandShakeReply.read(buf, key, &opts, self._compression_opts != null, stream); 176 162 errdefer self.close(.{ .code = 1001 }) catch unreachable; 177 163 178 164 // Set up compression with agreed-on parameters 179 - try self.setupCompression(res.compression); 165 + if (res.compression) { 166 + try self.setupCompression(); 167 + } 180 168 181 169 // We might have read more than handshake response. If so, readHandshakeReply 182 170 // has positioned the extra data at the start of the buffer, but we need ··· 184 172 self._reader.pos = res.over_read; 185 173 } 186 174 187 - fn setupCompression(self: *Client, agreed: ?ServerHandshake.Compression) !void { 188 - if (agreed == null) { 189 - self._compression_opts = null; 190 - } else if (self._compression_opts != null) { 191 - self._compression_opts.?.client_no_context_takeover = agreed.?.client_no_context_takeover; 192 - self._compression_opts.?.server_no_context_takeover = agreed.?.server_no_context_takeover; 193 - } else { 194 - unreachable; // HandShakeReply would return an error 195 - } 196 - if (self._compression_opts) |c| { 197 - self._reader.decompressor = .{}; 198 - self._reader.decompressor_reset = c.server_no_context_takeover; 199 - if (c.write_threshold == null) { 200 - return; 201 - } 202 - const allocator = self._reader.large_buffer_provider.allocator; 203 - const compression = try allocator.create(Compression); 204 - errdefer allocator.destroy(compression); 205 - compression.* = .{ 206 - .compressor = undefined, 207 - .write_treshold = c.write_threshold.?, 208 - .reset = c.server_no_context_takeover, 209 - .retain_writer = c.retain_write_buffer, 210 - .writer = std.ArrayList(u8).init(allocator), 211 - }; 212 - compression.compressor = try Compression.Type.init(compression.writer.writer(), .{}); 213 - self._compression = compression; 214 - } 175 + fn setupCompression(self: *Client) !void { 176 + std.debug.assert(self._compression_opts != null); 177 + self._reader.allow_compressed = true; 178 + 179 + const allocator = self._reader.large_buffer_provider.allocator; 180 + const config = self._compression_opts.?; 181 + self._compression = .{ 182 + .allocator = allocator, 183 + .write_treshold = config.write_threshold.?, 184 + .retain_writer = config.retain_write_buffer, 185 + .writer = std.Io.Writer.Allocating.init(allocator), 186 + }; 215 187 } 216 188 217 189 pub fn readLoop(self: *Client, handler: anytype) !void { ··· 367 339 } 368 340 369 341 pub fn writeFrame(self: *Client, op_code: proto.OpCode, data: []u8) !void { 370 - var payload = data; 371 - var compressed = false; 372 - if (self._compression) |c| { 373 - if (data.len >= c.write_treshold and (op_code == .binary or op_code == .text)) { 374 - compressed = true; 342 + const payload = data; 343 + const compressed = false; 344 + // if (self._compression) |c| { 345 + // if (data.len >= c.write_treshold and (op_code == .binary or op_code == .text)) { 346 + // compressed = true; 375 347 376 - var writer = &c.writer; 377 - var compressor = &c.compressor; 378 - var fbs = std.io.fixedBufferStream(data); 379 - _ = try compressor.compress(fbs.reader()); 380 - try compressor.flush(); 381 - payload = writer.items[0 .. writer.items.len - 4]; 348 + // var writer = &c.writer; 349 + // var compressor = &c.compressor; 350 + // var fbs = std.io.fixedBufferStream(data); 351 + // _ = try compressor.compress(fbs.reader()); 352 + // try compressor.flush(); 353 + // payload = writer.items[0 .. writer.items.len - 4]; 382 354 383 - if (c.reset) { 384 - c.compressor = try Compression.Type.init(writer.writer(), .{}); 385 - } 386 - } 387 - } 388 - defer if (compressed) { 389 - const c = self._compression.?; 390 - if (c.retain_writer) { 391 - c.compressor.wrt.context.clearRetainingCapacity(); 392 - } else { 393 - c.compressor.wrt.context.clearAndFree(); 394 - } 395 - }; 355 + // if (c.reset) { 356 + // c.compressor = try Compression.Type.init(writer.writer(), .{}); 357 + // } 358 + // } 359 + // } 360 + // defer if (compressed) { 361 + // const c = self._compression.?; 362 + // if (c.retain_writer) { 363 + // c.compressor.wrt.context.clearRetainingCapacity(); 364 + // } else { 365 + // c.compressor.wrt.context.clearAndFree(); 366 + // } 367 + // }; 396 368 397 369 // maximum possible prefix length. op_code + length_type + 8byte length + 4 byte mask 398 370 var buf: [14]u8 = undefined; ··· 423 395 // wraps a net.Stream and optional a tls.Client 424 396 pub const Stream = struct { 425 397 stream: net.Stream, 426 - tls_client: ?tls.Client = null, 398 + tls_client: ?*TLSClient = null, 427 399 428 - pub fn init(stream: net.Stream, tls_client: ?tls.Client) Stream { 400 + pub fn init(stream: net.Stream, tls_client: ?*TLSClient) Stream { 429 401 return .{ 430 402 .stream = stream, 431 403 .tls_client = tls_client, ··· 433 405 } 434 406 435 407 pub fn close(self: *Stream) void { 436 - if (self.tls_client) |*tls_client| { 437 - _ = tls_client.writeEnd(self.stream, "", true) catch {}; 408 + if (self.tls_client) |tls_client| { 409 + tls_client.deinit(); 438 410 } 439 411 440 412 // std.posix.close panics on EBADF ··· 457 429 } 458 430 459 431 pub fn read(self: *Stream, buf: []u8) !usize { 460 - if (self.tls_client) |*tls_client| { 461 - return tls_client.read(self.stream, buf); 432 + if (self.tls_client) |tls_client| { 433 + var w: std.Io.Writer = .fixed(buf); 434 + while (true) { 435 + const n = try tls_client.client.reader.stream(&w, .limited(buf.len)); 436 + if (n != 0) { 437 + return n; 438 + } 439 + } 462 440 } 463 441 return self.stream.read(buf); 464 442 } 465 443 466 444 pub fn writeAll(self: *Stream, data: []const u8) !void { 467 - if (self.tls_client) |*tls_client| { 468 - return tls_client.writeAll(self.stream, data); 445 + if (self.tls_client) |tls_client| { 446 + try tls_client.client.writer.writeAll(data); 447 + // I know this looks silly, but as far as I can tell, this is what 448 + // we need to do. 449 + try tls_client.client.writer.flush(); 450 + try tls_client.stream_writer.interface.flush(); 451 + return; 469 452 } 470 453 return self.stream.writeAll(data); 471 454 } ··· 496 479 } 497 480 }; 498 481 482 + const TLSClient = struct { 483 + client: tls.Client, 484 + stream: net.Stream, 485 + stream_writer: net.Stream.Writer, 486 + stream_reader: net.Stream.Reader, 487 + arena: std.heap.ArenaAllocator, 488 + 489 + fn init(allocator: Allocator, stream: net.Stream, config: *const Client.Config) !*TLSClient { 490 + var arena = std.heap.ArenaAllocator.init(allocator); 491 + errdefer arena.deinit(); 492 + 493 + const aa = arena.allocator(); 494 + 495 + const bundle = config.ca_bundle orelse blk: { 496 + var b = Bundle{}; 497 + try b.rescan(aa); 498 + break :blk b; 499 + }; 500 + 501 + // The TLS input and output have to be max_ciphertext_record_len each. 502 + // It isn't clear to me how big the un-encrypted reader and writer 503 + // need to be. I would think 0, but that will fail an assertion. I 504 + // don't think that it's right that we need 4 buffers, but apparently 505 + // we do. Until i figure this out, using 4 x max_ciphertext_record_len 506 + // seems like the only safe choice. 507 + const buf_len = std.crypto.tls.max_ciphertext_record_len; 508 + var buf = try aa.alloc(u8, buf_len * 4); 509 + 510 + const self = try aa.create(TLSClient); 511 + self.* = .{ 512 + .stream = stream, 513 + .arena = arena, 514 + .client = undefined, 515 + .stream_writer = stream.writer(buf.ptr[0..buf_len][0..buf_len]), 516 + .stream_reader = stream.reader(buf.ptr[buf_len .. 2 * buf_len][0..buf_len]), 517 + }; 518 + 519 + self.client = try tls.Client.init( 520 + self.stream_reader.interface(), 521 + &self.stream_writer.interface, 522 + .{ 523 + .ca = .{ .bundle = bundle }, 524 + .host = .{ .explicit = config.host }, 525 + .read_buffer = buf.ptr[2 * buf_len .. 3 * buf_len][0..buf_len], 526 + .write_buffer = buf.ptr[3 * buf_len .. 4 * buf_len][0..buf_len], 527 + }, 528 + ); 529 + 530 + return self; 531 + } 532 + 533 + fn deinit(self: *TLSClient) void { 534 + _ = self.client.end() catch {}; 535 + self.arena.deinit(); 536 + } 537 + }; 538 + 499 539 fn generateKey() [16]u8 { 500 540 if (comptime @import("builtin").is_test) { 501 541 return [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; ··· 511 551 return m; 512 552 } 513 553 514 - fn sendHandshake(path: []const u8, key: []const u8, buf: []u8, opts: *const Client.HandshakeOpts, compression: ?CompressionOpts, stream: anytype) !void { 554 + fn sendHandshake(path: []const u8, key: []const u8, buf: []u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !void { 515 555 @memcpy(buf[0..4], "GET "); 516 556 var pos: usize = 4; 517 557 var end = pos + path.len; ··· 531 571 @memcpy(buf[pos..end], key); 532 572 } 533 573 534 - if (compression) |c| { 535 - // TODO: also advertise no_context_takeover if false 574 + if (compression) { 536 575 // NOTE: client_max_window_bits is unsupported 537 - { 538 - const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate"; 539 - pos = end; 540 - end = pos + permessage_deflate.len; 541 - @memcpy(buf[pos..end], permessage_deflate); 542 - } 543 - if (c.server_no_context_takeover) { 544 - const server_no_context_takeover = "; server_no_context_takeover"; 545 - pos = end; 546 - end = pos + server_no_context_takeover.len; 547 - @memcpy(buf[pos..end], server_no_context_takeover); 548 - } 549 - if (c.client_no_context_takeover) { 550 - const client_no_context_takeover = "; client_no_context_takeover"; 551 - pos = end; 552 - end = pos + client_no_context_takeover.len; 553 - @memcpy(buf[pos..end], client_no_context_takeover); 554 - } 576 + const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover"; 577 + pos = end; 578 + end = pos + permessage_deflate.len; 579 + @memcpy(buf[pos..end], permessage_deflate); 555 580 } 556 581 557 582 { ··· 580 605 } 581 606 582 607 const HandShakeReply = struct { 583 - compression: ?ServerHandshake.Compression, 608 + compression: bool, 584 609 over_read: usize, 585 610 586 - fn read(buf: []u8, key: []const u8, opts: *const Client.HandshakeOpts, compression: ?CompressionOpts, stream: anytype) !HandShakeReply { 611 + fn read(buf: []u8, key: []const u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !HandShakeReply { 587 612 const timeout_ms = opts.timeout_ms; 588 613 const deadline = std.time.milliTimestamp() + timeout_ms; 589 614 try stream.readTimeout(timeout_ms); ··· 591 616 var pos: usize = 0; 592 617 var line_start: usize = 0; 593 618 var complete_response: u8 = 0; 594 - var server_compression: ?ServerHandshake.Compression = null; 619 + var server_compression: bool = false; 595 620 596 621 while (true) { 597 622 const n = stream.read(buf[pos..]) catch |err| switch (err) { ··· 612 637 std.mem.copyForwards(u8, buf[0..over_read], buf[line_start + 2 .. pos]); 613 638 try stream.readTimeout(0); 614 639 return .{ 615 - .compression = server_compression, 616 640 .over_read = over_read, 641 + .compression = server_compression, 617 642 }; 618 643 } 619 644 ··· 674 699 complete_response |= 8; 675 700 }, 676 701 24 => if (std.mem.eql(u8, line[0..i], "sec-websocket-extensions")) { 677 - server_compression = try parseExtension(line[i + 1 ..]); 678 - if (server_compression) |agreed| { 679 - if (compression) |offered| { 680 - if ((!agreed.client_no_context_takeover and offered.client_no_context_takeover) or 681 - (!agreed.server_no_context_takeover and offered.server_no_context_takeover)) 682 - { 683 - return error.InvalidExtensionHeader; 684 - } 685 - } else { 702 + if (try parseExtension(line[i + 1 ..])) |sc| { 703 + if (!compression) { 704 + // server is saying compression, but we didn't ask for it. 686 705 return error.InvalidExtensionHeader; 687 706 } 707 + if (!sc.client_no_context_takeover or !sc.server_no_context_takeover) { 708 + // as of Zig 0.15, we no longer support context takeover 709 + // We told the server this, it should have respected it. 710 + return error.InvalidExtensionHeader; 711 + } 712 + 713 + server_compression = true; 688 714 } 689 715 }, 690 716 else => {}, // some other header we don't care about ··· 954 980 ._closed = false, 955 981 ._own_bp = true, 956 982 ._mask_fn = generateMask, 983 + ._compression_opts = null, 957 984 .stream = .{ .stream = stream }, 958 985 ._reader = Reader.init(reader_buf, bp, null), 959 986 };
+33 -65
src/proto.zig
··· 87 87 // fragment), the state of the fragmented message is maintained here.) 88 88 fragment: ?Fragmented, 89 89 90 - decompressor: ?DecompressorType, 90 + allow_compressed: bool, 91 91 92 92 // if we returned a decompressed message, it's stored here so that we can 93 93 // cleanup when the user is done with the message 94 94 decompress_writer: ?buffer.Writer, 95 95 96 - // when client_no_context_takeover, we reset the decompressor after every 97 - // message 98 - decompressor_reset: bool, 99 - 100 - const DecompressorType = std.compress.flate.Decompressor(std.io.FixedBufferStream([]const u8).Reader); 96 + const DecompressorType = std.compress.flate.Decompress; 101 97 102 98 pub fn init(static: []u8, large_buffer_provider: *buffer.Provider, compression: ?Compression) Reader { 103 - var decompressor_reset = false; 104 - var decompressor: ?DecompressorType = null; 105 - if (compression) |c| { 106 - decompressor = .{}; 107 - decompressor_reset = c.client_no_context_takeover; 108 - } 109 - 110 99 return .{ 111 100 .pos = 0, 112 101 .start = 0, ··· 115 104 .message_len = 0, 116 105 .fragment = null, 117 106 .decompress_writer = null, 118 - .decompressor = decompressor, 119 - .decompressor_reset = decompressor_reset, 107 + .allow_compressed = compression != null, 120 108 .large_buffer_provider = large_buffer_provider, 121 109 }; 122 110 } 123 111 124 112 pub fn deinit(self: *Reader) void { 125 - if (self.fragment) |f| { 113 + if (self.fragment) |*f| { 126 114 f.deinit(); 127 115 } 128 116 if (self.decompress_writer) |*dw| { ··· 265 253 // FIN, RSV1, RSV2, RSV3, OP,OP,OP,OP 266 254 // RSV2 and RSV3 should never be set, and RSV1 should not be set 267 255 // when compression is disabled 268 - const rsv_bits: u8 = if (self.decompressor == null or is_continuation) 112 else 48; 256 + const rsv_bits: u8 = if (self.allow_compressed == false or is_continuation) 112 else 48; 269 257 if (byte1 & rsv_bits != 0) { 270 258 return error.ReservedFlags; 271 259 } 272 260 273 261 const compressed = byte1 & 64 == 64; 274 262 if (compressed) { 275 - if (self.decompressor == null) { 263 + if (self.allow_compressed == false) { 276 264 return error.CompressionDisabled; 277 265 } 278 266 } ··· 371 359 // call to "read" to do this, because we don't know when that'll be. 372 360 pub fn done(self: *Reader, message_type: Message.Type) void { 373 361 if (message_type == .text or message_type == .binary) { 374 - if (self.fragment) |f| { 362 + if (self.fragment) |*f| { 375 363 f.deinit(); 376 364 self.fragment = null; 377 365 } 378 366 if (self.decompress_writer) |*dw| { 379 367 dw.deinit(); 380 368 self.decompress_writer = null; 381 - 382 - if (self.decompressor_reset) { 383 - self.decompressor = .{}; 384 - } 385 369 } 386 370 } 387 371 ··· 420 404 fn decompress(self: *Reader, compressed: []const u8) ![]u8 { 421 405 const provider = self.large_buffer_provider; 422 406 407 + var dumb: [32]u8 = undefined; 423 408 var writer: buffer.Writer = undefined; 424 409 if (compressed.len < provider.pool_buffer_size) { 425 - writer = .{ 426 - .pooled = true, 427 - .provider = provider, 428 - .buf = try provider.pool.acquireOrCreate(), 429 - }; 410 + const buf = try provider.pool.acquireOrCreate(); 411 + writer = .init(buf, true, provider, &dumb); 430 412 } else { 431 - writer = .{ 432 - .pooled = false, 433 - .provider = provider, 434 - .buf = try provider.allocator.alloc(u8, @intFromFloat(@as(f64, @floatFromInt(compressed.len)) * 1.25)), 435 - }; 413 + const buf = try provider.allocator.alloc(u8, @intFromFloat(@as(f64, @floatFromInt(compressed.len)) * 1.25)); 414 + writer = .init(buf, false, provider, &dumb); 436 415 } 437 416 438 417 errdefer writer.deinit(); 439 418 440 - var decompressor = &self.decompressor.?; 441 - { 442 - var reader = std.io.fixedBufferStream(compressed); 443 - decompressor.setReader(reader.reader()); 444 - decompressor.decompress(&writer) catch |err| switch (err) { 445 - error.EndOfStream => {}, 446 - else => return error.CompressionError, 447 - }; 448 - } 449 - 450 - { 451 - var reader = std.io.fixedBufferStream(&[_]u8{ 0x00, 0x00, 0xff, 0xff }); 452 - decompressor.setReader(reader.reader()); 453 - decompressor.decompress(&writer) catch |err| switch (err) { 454 - error.EndOfStream => {}, 455 - else => return error.CompressionError, 456 - }; 457 - } 419 + var reader = std.Io.Reader.fixed(compressed); 420 + var decompressor = std.compress.flate.Decompress.init(&reader, .raw, &.{}); 421 + const n = decompressor.reader.streamRemaining(&writer.interface) catch { 422 + return error.CompressionError; 423 + }; 458 424 459 425 self.decompress_writer = writer; 460 - return writer.buf[0..writer.pos]; 426 + return writer.buf[0..n]; 461 427 } 462 428 463 429 inline fn usingLargeBuffer(self: *const Reader) bool { ··· 480 446 compressed: bool, 481 447 type: Message.Type, 482 448 buf: std.ArrayList(u8), 449 + allocator: std.mem.Allocator, 483 450 484 451 pub fn init(bp: *buffer.Provider, compressed: bool, message_type: Message.Type, value: []const u8) !Fragmented { 485 - var buf = std.ArrayList(u8).init(bp.allocator); 486 - try buf.ensureTotalCapacity(value.len * 2); 452 + var buf: std.ArrayList(u8) = .empty; 453 + try buf.ensureTotalCapacity(bp.allocator, value.len * 2); 487 454 buf.appendSliceAssumeCapacity(value); 488 455 489 456 return .{ ··· 491 458 .type = message_type, 492 459 .compressed = compressed, 493 460 .max = bp.max_buffer_size, 461 + .allocator = bp.allocator, 494 462 }; 495 463 } 496 464 497 - pub fn deinit(self: Fragmented) void { 498 - self.buf.deinit(); 465 + pub fn deinit(self: *Fragmented) void { 466 + self.buf.deinit(self.allocator); 499 467 } 500 468 501 469 pub fn add(self: *Fragmented, value: []const u8) !void { 502 470 if (self.buf.items.len + value.len > self.max) { 503 471 return error.TooLarge; 504 472 } 505 - try self.buf.appendSlice(value); 473 + try self.buf.appendSlice(self.allocator, value); 506 474 } 507 475 508 476 // Optimization so that we don't over-allocate on our last frame. ··· 511 479 if (total_len > self.max) { 512 480 return error.TooLarge; 513 481 } 514 - try self.buf.ensureTotalCapacityPrecise(total_len); 482 + try self.buf.ensureTotalCapacityPrecise(self.allocator, total_len); 515 483 self.buf.appendSliceAssumeCapacity(value); 516 484 return self.buf.items; 517 485 } ··· 680 648 681 649 var is_fragmented = false; 682 650 var fragment_count: usize = 0; 683 - var fragment = std.ArrayList(u8).init(arena); 651 + var fragment: std.ArrayList(u8) = .empty; 684 652 685 653 var i: usize = 0; 686 654 while (i < MESSAGE_TO_SEND) { ··· 710 678 } 711 679 fragment_count += 1; 712 680 713 - try fragment.appendSlice(try arena.dupe(u8, buf)); 681 + try fragment.appendSlice(arena, try arena.dupe(u8, buf)); 714 682 715 683 if (is_fin) { 716 684 // this was the last message in our fragment ··· 818 786 var f = try Fragmented.init(&bp, false, .binary, payload); 819 787 defer f.deinit(); 820 788 821 - var expected = std.ArrayList(u8).init(t.allocator); 822 - defer expected.deinit(); 823 - try expected.appendSlice(payload); 789 + var expected: std.ArrayList(u8) = .empty; 790 + defer expected.deinit(t.allocator); 791 + try expected.appendSlice(t.allocator, payload); 824 792 825 793 const number_of_adds = random.uintAtMost(usize, 30); 826 794 for (0..number_of_adds) |_| { 827 795 payload = buf[0 .. random.uintAtMost(usize, 99) + 1]; 828 796 random.bytes(payload); 829 797 try f.add(payload); 830 - try expected.appendSlice(payload); 798 + try expected.appendSlice(t.allocator, payload); 831 799 } 832 800 payload = buf[0 .. random.uintAtMost(usize, 99) + 1]; 833 801 random.bytes(payload); 834 - try expected.appendSlice(payload); 802 + try expected.appendSlice(t.allocator, payload); 835 803 836 804 try t.expectString(expected.items, try f.last(payload)); 837 805 }
+99 -135
src/server/handshake.zig
··· 122 122 }; 123 123 } 124 124 125 - pub fn createReply(key: []const u8, headers: *const KeyValue, compression: ?websocket.Compression, buf: []u8) ![]const u8 { 125 + pub fn createReply(key: []const u8, headers_: ?*KeyValue, compression: bool, buf: []u8) ![]const u8 { 126 126 const HEADER = 127 127 "HTTP/1.1 101 Switching Protocols\r\n" ++ 128 128 "Upgrade: websocket\r\n" ++ ··· 145 145 pos = end; 146 146 } 147 147 148 - if (compression) |c| { 149 - { 150 - const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate"; 151 - const end = pos + permessage_deflate.len; 152 - @memcpy(buf[pos..end], permessage_deflate); 153 - pos = end; 154 - } 155 - if (c.server_no_context_takeover) { 156 - const server_no_context_takeover = "; server_no_context_takeover"; 157 - const end = pos + server_no_context_takeover.len; 158 - @memcpy(buf[pos..end], server_no_context_takeover); 159 - pos = end; 160 - } 161 - if (c.client_no_context_takeover) { 162 - const client_no_context_takeover = "; client_no_context_takeover"; 163 - const end = pos + client_no_context_takeover.len; 164 - @memcpy(buf[pos..end], client_no_context_takeover); 165 - pos = end; 166 - } 148 + if (compression) { 149 + const permessage_deflate = 150 + "\r\nSec-WebSocket-Extensions: permessage-deflate" ++ 151 + "; server_no_context_takeover" ++ 152 + "; client_no_context_takeover"; 153 + 154 + const end = pos + permessage_deflate.len; 155 + @memcpy(buf[pos..end], permessage_deflate); 156 + pos = end; 167 157 } 168 158 169 - for (headers.keys[0..headers.len], headers.values[0..headers.len]) |k, v| { 170 - pos += (try std.fmt.bufPrint(buf[pos..], "\r\n{s}: {s}", .{ k, v })).len; 159 + if (headers_) |headers| { 160 + for (headers.keys[0..headers.len], headers.values[0..headers.len]) |k, v| { 161 + pos += (try std.fmt.bufPrint(buf[pos..], "\r\n{s}: {s}", .{ k, v })).len; 162 + } 171 163 } 172 164 173 165 const end = pos + 4; ··· 242 234 const buf = try allocator.alloc(u8, pool.buffer_size); 243 235 errdefer allocator.free(buf); 244 236 245 - const req_headers = try KeyValue.init(allocator, pool.max_req_headers); 237 + const req_headers = try Handshake.KeyValue.init(allocator, pool.max_req_headers); 246 238 errdefer req_headers.deinit(allocator); 247 239 248 - const res_headers = try KeyValue.init(allocator, pool.max_res_headers); 240 + const res_headers = try Handshake.KeyValue.init(allocator, pool.max_res_headers); 249 241 errdefer res_headers.deinit(allocator); 250 242 251 243 return .{ ··· 270 262 self.pool.release(self); 271 263 } 272 264 }; 273 - }; 274 265 275 - pub const KeyValue = struct { 276 - len: usize, 277 - keys: [][]const u8, 278 - values: [][]const u8, 266 + pub const KeyValue = struct { 267 + len: usize, 268 + keys: [][]const u8, 269 + values: [][]const u8, 270 + 271 + fn init(allocator: Allocator, max: usize) !KeyValue { 272 + const keys = try allocator.alloc([]const u8, max); 273 + errdefer allocator.free(keys); 279 274 280 - fn init(allocator: Allocator, max: usize) !KeyValue { 281 - const keys = try allocator.alloc([]const u8, max); 282 - errdefer allocator.free(keys); 275 + const values = try allocator.alloc([]const u8, max); 276 + errdefer allocator.free(values); 283 277 284 - const values = try allocator.alloc([]const u8, max); 285 - errdefer allocator.free(values); 278 + return .{ 279 + .len = 0, 280 + .keys = keys, 281 + .values = values, 282 + }; 283 + } 286 284 287 - return .{ 288 - .len = 0, 289 - .keys = keys, 290 - .values = values, 291 - }; 292 - } 285 + fn deinit(self: *const KeyValue, allocator: Allocator) void { 286 + allocator.free(self.keys); 287 + allocator.free(self.values); 288 + } 293 289 294 - fn deinit(self: *const KeyValue, allocator: Allocator) void { 295 - allocator.free(self.keys); 296 - allocator.free(self.values); 297 - } 290 + pub fn add(self: *KeyValue, key: []const u8, value: []const u8) void { 291 + const len = self.len; 292 + var keys = self.keys; 293 + if (len == keys.len) { 294 + return; 295 + } 298 296 299 - pub fn add(self: *KeyValue, key: []const u8, value: []const u8) void { 300 - const len = self.len; 301 - var keys = self.keys; 302 - if (len == keys.len) { 303 - return; 297 + keys[len] = key; 298 + self.values[len] = value; 299 + self.len = len + 1; 304 300 } 305 301 306 - keys[len] = key; 307 - self.values[len] = value; 308 - self.len = len + 1; 309 - } 310 - 311 - pub fn get(self: *const KeyValue, needle: []const u8) ?[]const u8 { 312 - const keys = self.keys[0..self.len]; 313 - loop: for (keys, 0..) |key, i| { 314 - // This is largely a reminder to myself that std.mem.eql isn't 315 - // particularly fast. Here we at least avoid the 1 extra ptr 316 - // equality check that std.mem.eql does, but we could do better 317 - // TODO: monitor https://github.com/ziglang/zig/issues/8689 318 - if (needle.len != key.len) { 319 - continue; 320 - } 321 - for (needle, key) |n, k| { 322 - if (n != k) { 323 - continue :loop; 302 + pub fn get(self: *const KeyValue, needle: []const u8) ?[]const u8 { 303 + const keys = self.keys[0..self.len]; 304 + loop: for (keys, 0..) |key, i| { 305 + // This is largely a reminder to myself that std.mem.eql isn't 306 + // particularly fast. Here we at least avoid the 1 extra ptr 307 + // equality check that std.mem.eql does, but we could do better 308 + // TODO: monitor https://github.com/ziglang/zig/issues/8689 309 + if (needle.len != key.len) { 310 + continue; 324 311 } 312 + for (needle, key) |n, k| { 313 + if (n != k) { 314 + continue :loop; 315 + } 316 + } 317 + return self.values[i]; 325 318 } 326 - return self.values[i]; 319 + 320 + return null; 327 321 } 328 322 329 - return null; 330 - } 323 + pub fn iterator(self: *const KeyValue) Iterator { 324 + const len = self.len; 325 + return .{ 326 + .pos = 0, 327 + .keys = self.keys[0..len], 328 + .values = self.values[0..len], 329 + }; 330 + } 331 331 332 - pub fn iterator(self: *const KeyValue) Iterator { 333 - const len = self.len; 334 - return .{ 335 - .pos = 0, 336 - .keys = self.keys[0..len], 337 - .values = self.values[0..len], 338 - }; 339 - } 332 + pub const Iterator = struct { 333 + pos: usize, 334 + keys: [][]const u8, 335 + values: [][]const u8, 340 336 341 - pub const Iterator = struct { 342 - pos: usize, 343 - keys: [][]const u8, 344 - values: [][]const u8, 337 + const KV = struct { 338 + key: []const u8, 339 + value: []const u8, 340 + }; 345 341 346 - const KV = struct { 347 - key: []const u8, 348 - value: []const u8, 349 - }; 342 + pub fn next(self: *Iterator) ?KV { 343 + const pos = self.pos; 344 + if (pos == self.keys.len) { 345 + return null; 346 + } 350 347 351 - pub fn next(self: *Iterator) ?KV { 352 - const pos = self.pos; 353 - if (pos == self.keys.len) { 354 - return null; 348 + self.pos = pos + 1; 349 + return .{ 350 + .key = self.keys[pos], 351 + .value = self.values[pos], 352 + }; 355 353 } 356 - 357 - self.pos = pos + 1; 358 - return .{ 359 - .key = self.keys[pos], 360 - .value = self.values[pos], 361 - }; 362 - } 354 + }; 363 355 }; 364 356 }; 365 357 ··· 525 517 526 518 test "handshake: reply" { 527 519 var buf: [512]u8 = undefined; 528 - var res_headers = try KeyValue.init(t.allocator, 2); 529 - defer res_headers.deinit(t.allocator); 530 520 531 521 { 532 522 // no compression ··· 535 525 "Upgrade: websocket\r\n" ++ 536 526 "Connection: upgrade\r\n" ++ 537 527 "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n\r\n"; 538 - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, null, &buf)); 539 - } 540 - 541 - { 542 - // compression 543 - const expected = 544 - "HTTP/1.1 101 Switching Protocols\r\n" ++ 545 - "Upgrade: websocket\r\n" ++ 546 - "Connection: upgrade\r\n" ++ 547 - "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ 548 - "Sec-WebSocket-Extensions: permessage-deflate\r\n\r\n"; 549 - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{}, &buf)); 528 + try t.expectString(expected, try Handshake.createReply("this is my key", null, false, &buf)); 550 529 } 551 530 552 531 { ··· 557 536 "Connection: upgrade\r\n" ++ 558 537 "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ 559 538 "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n\r\n"; 560 - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{ 561 - .client_no_context_takeover = true, 562 - .server_no_context_takeover = true, 563 - }, &buf)); 539 + try t.expectString(expected, try Handshake.createReply("this is my key", null, true, &buf)); 564 540 } 565 541 566 542 // With custom headers 543 + var res_headers = try Handshake.KeyValue.init(t.allocator, 2); 544 + defer res_headers.deinit(t.allocator); 567 545 res_headers.add("Set-Cookie", "Yummy!"); 546 + 568 547 { 569 548 // no compression 570 549 const expected = ··· 573 552 "Connection: upgrade\r\n" ++ 574 553 "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ 575 554 "Set-Cookie: Yummy!\r\n\r\n"; 576 - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, null, &buf)); 577 - } 578 - 579 - { 580 - // compression 581 - const expected = 582 - "HTTP/1.1 101 Switching Protocols\r\n" ++ 583 - "Upgrade: websocket\r\n" ++ 584 - "Connection: upgrade\r\n" ++ 585 - "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ 586 - "Sec-WebSocket-Extensions: permessage-deflate\r\n" ++ 587 - "Set-Cookie: Yummy!\r\n\r\n"; 588 - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{}, &buf)); 555 + try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, false, &buf)); 589 556 } 590 557 591 558 { ··· 597 564 "Sec-Websocket-Accept: flzHu2DevQ2dSCSVqKSii5e9C2o=\r\n" ++ 598 565 "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover\r\n" ++ 599 566 "Set-Cookie: Yummy!\r\n\r\n"; 600 - try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, .{ 601 - .client_no_context_takeover = true, 602 - .server_no_context_takeover = true, 603 - }, &buf)); 567 + try t.expectString(expected, try Handshake.createReply("this is my key", &res_headers, true, &buf)); 604 568 } 605 569 } 606 570 607 571 test "KeyValue: get" { 608 572 const allocator = t.allocator; 609 - var kv = try KeyValue.init(allocator, 2); 573 + var kv = try Handshake.KeyValue.init(allocator, 2); 610 574 defer kv.deinit(t.allocator); 611 575 612 576 var key = "content-type".*; ··· 621 585 } 622 586 623 587 test "KeyValue: ignores beyond max" { 624 - var kv = try KeyValue.init(t.allocator, 2); 588 + var kv = try Handshake.KeyValue.init(t.allocator, 2); 625 589 defer kv.deinit(t.allocator); 626 590 627 591 var n1 = "content-length".*; ··· 696 660 var hs = p.acquire() catch unreachable; 697 661 std.debug.assert(hs.buf[0] == 0); 698 662 hs.buf[0] = 255; 699 - std.time.sleep(random.uintAtMost(u32, 100000)); 663 + std.Thread.sleep(random.uintAtMost(u32, 100000)); 700 664 hs.buf[0] = 0; 701 665 p.release(hs); 702 666 }
+113 -119
src/server/server.zig
··· 110 110 } 111 111 } 112 112 113 + if (config.compression != null) { 114 + log.err("Compression is disabled as part of the 0.15 upgrade. I do hope to re-enable it soon.", .{}); 115 + return error.InvalidConfiguraion; 116 + } 117 + 113 118 const signals = try allocator.alloc(posix.fd_t, config.workerCount()); 114 119 errdefer allocator.free(signals); 115 120 ··· 253 258 started += 1; 254 259 } 255 260 256 - log.info("starting nonblocking worker to listen on {}", .{address}); 261 + log.info("starting nonblocking worker to listen on {f}", .{address}); 257 262 258 263 // in case startInNewThread is waiting 259 264 self._cond.signal(); ··· 344 349 log.err("failed to accept socket: {}", .{err}); 345 350 continue; 346 351 }; 347 - log.debug("({}) connected", .{address}); 352 + log.debug("({f}) connected", .{address}); 348 353 349 354 const thread = std.Thread.spawn(.{}, Self.handleConnection, .{ self, socket, address, ctx }) catch |err| { 350 355 posix.close(socket); 351 - log.err("({}) failed to spawn connection thread: {}", .{ address, err }); 356 + log.err("({f}) failed to spawn connection thread: {}", .{ address, err }); 352 357 continue; 353 358 }; 354 359 thread.detach(); ··· 359 364 // Wrapper around _handleConnection so that we can handle erros 360 365 fn handleConnection(self: *Self, socket: posix.socket_t, address: net.Address, ctx: anytype) void { 361 366 self._handleConnection(socket, address, ctx) catch |err| { 362 - log.err("({}) uncaught error in connection handler: {}", .{ address, err }); 367 + log.err("({f}) uncaught error in connection handler: {}", .{ address, err }); 363 368 }; 364 369 } 365 370 ··· 382 387 } 383 388 if (hc.handler != null) { 384 389 // if we have a handler, the our handshake completed 385 - try conn_manager.setupCompression(hc, compression); 390 + if (compression) { 391 + try conn_manager.setupCompression(hc); 392 + } 386 393 break; 387 394 } 388 395 if (timestamp() > deadline) { ··· 430 437 if (conn_manager.count() == 0) { 431 438 return; 432 439 } 433 - std.time.sleep(std.time.ns_per_ms * 100); 440 + std.Thread.sleep(std.time.ns_per_ms * 100); 434 441 } 435 442 } 436 443 ··· 520 527 521 528 var it = self.loop.wait(timeout) catch |err| { 522 529 log.err("failed to wait on events: {}", .{err}); 523 - std.time.sleep(std.time.ns_per_s); 530 + std.Thread.sleep(std.time.ns_per_s); 524 531 continue; 525 532 }; 526 533 ··· 530 537 if (data == 0) { 531 538 self.accept(listener, now) catch |err| { 532 539 log.err("accept error: {}", .{err}); 533 - std.time.sleep(std.time.ns_per_ms); 540 + std.Thread.sleep(std.time.ns_per_ms); 534 541 }; 535 542 continue; 536 543 } ··· 581 588 // this connection has timed out. Don't use self.cleanup since there's 582 589 // a bunch of stuff we can assume here..like there's no handler or reader 583 590 conn.closeSocket(); 584 - log.debug("({}) handshake timeout", .{conn.address}); 591 + log.debug("({f}) handshake timeout", .{conn.address}); 585 592 if (hc.handshake) |h| { 586 593 h.release(); 587 594 } ··· 606 613 return if (err == error.WouldBlock) {} else err; 607 614 }; 608 615 609 - log.debug("({}) connected", .{address}); 616 + log.debug("({f}) connected", .{address}); 610 617 611 618 { 612 619 errdefer posix.close(socket); ··· 642 649 var success = false; 643 650 if (hc.handler == null) { 644 651 success = self.dataForHandshake(hc) catch |err| blk: { 645 - log.err("({}) error processing handshake: {}", .{ hc.conn.address, err }); 652 + log.err("({f}) error processing handshake: {}", .{ hc.conn.address, err }); 646 653 break :blk false; 647 654 }; 648 655 } else { ··· 662 669 self.base.cleanupConn(hc); 663 670 } else { 664 671 self.loop.monitorRead(hc, true) catch |err| { 665 - log.debug("({}) failed to add read event monitor: {}", .{ conn.address, err }); 672 + log.debug("({f}) failed to add read event monitor: {}", .{ conn.address, err }); 666 673 conn.closeSocket(); 667 674 self.base.cleanupConn(hc); 668 675 }; ··· 681 688 conn_manager.inactive(hc); 682 689 } 683 690 684 - try conn_manager.setupCompression(hc, compression); 691 + if (compression) { 692 + try conn_manager.setupCompression(hc); 693 + } 685 694 return true; 686 695 } 687 696 }; ··· 740 749 741 750 pub fn dataAvailable(self: *Self, hc: *HandlerConn(H), thread_buf: []u8) bool { 742 751 return self._dataAvailable(hc, thread_buf) catch |err| { 743 - log.err("({}) error processing client message: {}", .{ hc.conn.address, err }); 752 + log.err("({f}) error processing client message: {}", .{ hc.conn.address, err }); 744 753 return false; 745 754 }; 746 755 } ··· 1002 1011 return self.worker.conn_manager.compression != null; 1003 1012 } 1004 1013 1005 - pub fn setupConnection(self: *Self, hc: *HandlerConn(H), agreed: ?Compression) !void { 1006 - return self.worker.conn_manager.setupCompression(hc, agreed); 1014 + pub fn setupConnection( 1015 + self: *Self, 1016 + hc: *HandlerConn(H), 1017 + ) !void { 1018 + return self.worker.conn_manager.setupCompression(hc); 1007 1019 } 1008 1020 1009 1021 pub fn shutdown(self: *Self) void { ··· 1223 1235 self.lock.unlock(); 1224 1236 } 1225 1237 1226 - fn setupCompression(self: *Self, hc: *HandlerConn(H), agreed_: ?Compression) !void { 1227 - const agreed = agreed_ orelse { 1228 - return; 1229 - }; 1238 + fn setupCompression(self: *Self, hc: *HandlerConn(H)) !void { 1239 + const config = self.compression orelse return; 1230 1240 1231 - const configured = self.compression.?; 1232 - const merged = Compression{ 1233 - .write_threshold = configured.write_threshold, 1234 - .retain_write_buffer = configured.retain_write_buffer, 1235 - .client_no_context_takeover = agreed.client_no_context_takeover, 1236 - .server_no_context_takeover = agreed.server_no_context_takeover, 1237 - }; 1238 - hc.compression = merged; 1241 + hc.compression = config; 1239 1242 1240 - if (merged.write_threshold == null) { 1241 - // and we have a write threshold, we need to setup our 1242 - // connection's compression (read compression is configured in our reader) 1243 + if (config.write_threshold == null) { 1244 + // if write_treshold is null, then we never want to compress 1245 + // outgoing messages. We don't need to set the conn.compression 1246 + // field. 1247 + // We'll still [potentially] decompress incoming messages, but 1248 + // that's set on the proto. 1243 1249 return; 1244 1250 } 1245 1251 1246 - var compression = try self.compression_pool.create(); 1252 + const compression = try self.compression_pool.create(); 1247 1253 errdefer self.compression_pool.destroy(compression); 1248 1254 1249 1255 compression.* = .{ 1250 - .compressor = undefined, 1251 - .write_treshold = merged.write_threshold.?, 1252 - .reset = merged.server_no_context_takeover, 1253 - .retain_writer = merged.retain_write_buffer, 1254 - .writer = std.ArrayList(u8).init(self.allocator), 1256 + .allocator = self.allocator, 1257 + .write_treshold = config.write_threshold.?, 1258 + .retain_writer = config.retain_write_buffer, 1259 + .writer = std.Io.Writer.Allocating.init(self.allocator), 1255 1260 }; 1256 - compression.compressor = try Conn.Compression.Type.init(compression.writer.writer(), .{}); 1257 1261 hc.conn.compression = compression; 1258 1262 } 1259 1263 ··· 1299 1303 compression: ?*Conn.Compression = null, 1300 1304 1301 1305 const Compression = struct { 1302 - reset: bool, 1306 + allocator: Allocator, 1303 1307 retain_writer: bool, 1304 1308 write_treshold: usize, 1305 - compressor: Type, 1306 - writer: std.ArrayList(u8), 1307 - 1308 - const Type = std.compress.flate.Compressor(std.ArrayList(u8).Writer); 1309 + writer: std.Io.Writer.Allocating, 1309 1310 }; 1310 1311 1311 1312 pub fn isClosed(self: *Conn) bool { ··· 1371 1372 } 1372 1373 1373 1374 pub fn writeFrame(self: *Conn, op_code: OpCode, data: []const u8) !void { 1374 - var payload = data; 1375 - var compressed = false; 1376 - if (self.compression) |c| { 1377 - if (data.len >= c.write_treshold) { 1378 - compressed = true; 1375 + const payload = data; 1379 1376 1380 - var writer = &c.writer; 1381 - var compressor = &c.compressor; 1382 - var fbs = std.io.fixedBufferStream(data); 1383 - _ = try compressor.compress(fbs.reader()); 1384 - try compressor.flush(); 1385 - payload = writer.items[0 .. writer.items.len - 4]; 1377 + // Zig 0.15 compression disabled 1378 + const compressed = false; 1379 + // if (self.compression) |c| { 1380 + // if (data.len >= c.write_treshold) { 1381 + // compressed = true; 1382 + // var compressor = std.compress.flate.Compress.init(&c.writer.writer, &.{}, .{}); 1383 + // try compressor.writer.writeAll(data); 1384 + // try compressor.writer.flush(); 1385 + // const all = c.writer.written(); 1386 + // payload = all[0 .. all.len - 4]; 1387 + // } 1388 + // } 1386 1389 1387 - if (c.reset) { 1388 - c.compressor = try Conn.Compression.Type.init(writer.writer(), .{}); 1389 - } 1390 - } 1391 - } 1392 - defer if (compressed) { 1393 - const c = self.compression.?; 1394 - if (c.retain_writer) { 1395 - c.compressor.wrt.context.clearRetainingCapacity(); 1396 - } else { 1397 - c.compressor.wrt.context.clearAndFree(); 1398 - } 1399 - }; 1390 + // defer if (compressed) { 1391 + // const c = self.compression.?; 1392 + // if (c.retain_writer) { 1393 + // c.writer.clearRetainingCapacity(); 1394 + // } else { 1395 + // c.writer.deinit(); 1396 + // c.writer = std.Io.Writer.Allocating.init(c.allocator); 1397 + // } 1398 + // }; 1400 1399 1401 1400 // maximum possible prefix length. op_code + length_type + 8byte length 1402 1401 var buf: [10]u8 = undefined; ··· 1447 1446 pub fn writeBuffer(self: *Conn, allocator: Allocator, op_code: OpCode) Writer { 1448 1447 return .{ 1449 1448 .conn = self, 1449 + .buf = .empty, 1450 1450 .op_code = op_code, 1451 - .buf = std.ArrayList(u8).init(allocator), 1451 + .allocator = allocator, 1452 + .interface = .{ 1453 + .vtable = &.{ .drain = Writer.drain }, 1454 + .buffer = &.{}, 1455 + }, 1452 1456 }; 1453 1457 } 1454 1458 ··· 1461 1465 pub const Writer = struct { 1462 1466 conn: *Conn, 1463 1467 op_code: OpCode, 1464 - buf: std.ArrayList(u8), 1468 + allocator: Allocator, 1469 + buf: std.ArrayListUnmanaged(u8), 1470 + interface: std.io.Writer, 1465 1471 1466 1472 pub const Error = Allocator.Error; 1467 - pub const IOWriter = std.io.Writer(*Writer, error{OutOfMemory}, Writer.write); 1468 1473 1469 1474 pub fn deinit(self: *Writer) void { 1470 - self.buf.deinit(); 1471 - } 1472 - 1473 - pub fn writer(self: *Writer) IOWriter { 1474 - return .{ .context = self }; 1475 + self.buf.deinit(self.allocator); 1475 1476 } 1476 1477 1477 - pub fn write(self: *Writer, data: []const u8) Allocator.Error!usize { 1478 - try self.buf.appendSlice(data); 1479 - return data.len; 1478 + pub fn drain(io_w: *std.io.Writer, data: []const []const u8, splat: usize) error{WriteFailed}!usize { 1479 + _ = splat; 1480 + const self: *Writer = @alignCast(@fieldParentPtr("interface", io_w)); 1481 + self.buf.appendSlice(self.allocator, data[0]) catch return error.WriteFailed; 1482 + return data[0].len; 1480 1483 } 1481 1484 1482 - pub fn flush(self: *Writer) !void { 1483 - try self.conn.writeFrame(self.op_code, self.buf.items); 1485 + pub fn send(self: *Writer) !void { 1486 + return self.conn.writeFrame(self.op_code, self.buf.items) catch error.WriteFailed; 1484 1487 } 1485 1488 }; 1486 1489 }; 1487 1490 1488 - fn handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) struct { ?Compression, bool } { 1491 + fn handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) struct { bool, bool } { 1489 1492 return _handleHandshake(H, worker, hc, ctx) catch |err| { 1490 - log.warn("({}) uncaugh error processing handshake: {}", .{ hc.conn.address, err }); 1491 - return .{ null, false }; 1493 + log.warn("({f}) uncaugh error processing handshake: {}", .{ hc.conn.address, err }); 1494 + return .{ false, false }; 1492 1495 }; 1493 1496 } 1494 1497 1495 - fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) !struct { ?Compression, bool } { 1498 + fn _handleHandshake(comptime H: type, worker: anytype, hc: *HandlerConn(H), ctx: anytype) !struct { bool, bool } { 1496 1499 std.debug.assert(hc.handler == null); 1497 1500 1498 1501 var state = hc.handshake orelse blk: { ··· 1506 1509 const len = state.len; 1507 1510 1508 1511 if (len == buf.len) { 1509 - log.warn("({}) handshake request exceeded maximum configured size ({d})", .{ conn.address, buf.len }); 1510 - return .{ null, false }; 1512 + log.warn("({f}) handshake request exceeded maximum configured size ({d})", .{ conn.address, buf.len }); 1513 + return .{ false, false }; 1511 1514 } 1512 1515 1513 1516 const n = posix.read(hc.socket, buf[len..]) catch |err| { 1514 1517 switch (err) { 1515 - error.BrokenPipe, error.ConnectionResetByPeer => log.debug("({}) handshake connection closed: {}", .{ conn.address, err }), 1518 + error.BrokenPipe, error.ConnectionResetByPeer => log.debug("({f}) handshake connection closed: {}", .{ conn.address, err }), 1516 1519 error.WouldBlock => { 1517 1520 std.debug.assert(blockingMode()); 1518 - log.debug("({}) handshake timeout", .{conn.address}); 1521 + log.debug("({f}) handshake timeout", .{conn.address}); 1519 1522 }, 1520 - else => log.warn("({}) handshake error reading from socket: {}", .{ conn.address, err }), 1523 + else => log.warn("({f}) handshake error reading from socket: {}", .{ conn.address, err }), 1521 1524 } 1522 - return .{ null, false }; 1525 + return .{ false, false }; 1523 1526 }; 1524 1527 1525 1528 if (n == 0) { 1526 - log.debug("({}) handshake connection closed", .{conn.address}); 1527 - return .{ null, false }; 1529 + log.debug("({f}) handshake connection closed", .{conn.address}); 1530 + return .{ false, false }; 1528 1531 } 1529 1532 1530 1533 state.len = len + n; 1531 1534 var handshake = Handshake.parse(state) catch |err| { 1532 - log.debug("({}) error parsing handshake: {}", .{ conn.address, err }); 1535 + log.debug("({f}) error parsing handshake: {}", .{ conn.address, err }); 1533 1536 respondToHandshakeError(conn, err); 1534 - return .{ null, false }; 1537 + return .{ false, false }; 1535 1538 } orelse { 1536 1539 // we need more data 1537 - return .{ null, true }; 1540 + return .{ false, true }; 1538 1541 }; 1539 1542 1540 - var agreed_compression: ?Compression = null; 1541 - if (worker.compression) |configured_compression| { 1542 - if (handshake.compression) |request_compression| { 1543 - agreed_compression = .{ 1544 - .client_no_context_takeover = configured_compression.client_no_context_takeover or request_compression.client_no_context_takeover, 1545 - .server_no_context_takeover = configured_compression.server_no_context_takeover or request_compression.server_no_context_takeover, 1546 - }; 1547 - } 1548 - } 1549 - 1543 + const compression = handshake.compression != null and worker.compression != null; 1550 1544 defer state.release(); 1551 1545 hc.handshake = null; 1552 1546 ··· 1559 1553 } else { 1560 1554 respondToHandshakeError(conn, err); 1561 1555 } 1562 - log.debug("({}) " ++ @typeName(H) ++ ".init rejected request {}", .{ conn.address, err }); 1563 - return .{ null, false }; 1556 + log.debug("({f}) " ++ @typeName(H) ++ ".init rejected request {}", .{ conn.address, err }); 1557 + return .{ false, false }; 1564 1558 }; 1565 1559 1566 1560 hc.handler = handler; 1567 1561 1568 1562 var reply_buf: [2048]u8 = undefined; 1569 - const handshake_reply = try Handshake.createReply(handshake.key, handshake.res_headers, agreed_compression, &reply_buf); 1563 + const handshake_reply = try Handshake.createReply(handshake.key, handshake.res_headers, compression, &reply_buf); 1570 1564 try conn.writeFramed(handshake_reply); 1571 1565 1572 1566 if (comptime std.meta.hasFn(H, "afterInit")) { 1573 1567 const params = @typeInfo(@TypeOf(H.afterInit)).@"fn".params; 1574 1568 const res = if (params.len == 1) hc.handler.?.afterInit() else hc.handler.?.afterInit(ctx); 1575 1569 res catch |err| { 1576 - log.debug("({}) " ++ @typeName(H) ++ ".afterInit error: {}", .{ conn.address, err }); 1577 - return .{ null, false }; 1570 + log.debug("({f}) " ++ @typeName(H) ++ ".afterInit error: {}", .{ conn.address, err }); 1571 + return .{ false, false }; 1578 1572 }; 1579 1573 } 1580 1574 1581 - log.debug("({}) connection successfully upgraded", .{conn.address}); 1582 - return .{ agreed_compression, true }; 1575 + log.debug("({f}) connection successfully upgraded", .{conn.address}); 1576 + return .{ compression, true }; 1583 1577 } 1584 1578 1585 1579 fn handleClientData(comptime H: type, hc: *HandlerConn(H), allocator: Allocator, fba: *FixedBufferAllocator) bool { 1586 1580 std.debug.assert(hc.handshake == null); 1587 1581 return _handleClientData(H, hc, allocator, fba) catch |err| { 1588 - log.warn("({}) uncaugh error handling incoming data: {}", .{ hc.conn.address, err }); 1582 + log.warn("({f}) uncaugh error handling incoming data: {}", .{ hc.conn.address, err }); 1589 1583 return false; 1590 1584 }; 1591 1585 } ··· 1595 1589 var reader = &hc.reader.?; 1596 1590 reader.fill(conn.stream) catch |err| { 1597 1591 switch (err) { 1598 - error.BrokenPipe, error.Closed, error.ConnectionResetByPeer => log.debug("({}) connection closed: {}", .{ conn.address, err }), 1599 - else => log.warn("({}) error reading from connection: {}", .{ conn.address, err }), 1592 + error.BrokenPipe, error.Closed, error.ConnectionResetByPeer => log.debug("({f}) connection closed: {}", .{ conn.address, err }), 1593 + else => log.warn("({f}) error reading from connection: {}", .{ conn.address, err }), 1600 1594 } 1601 1595 return false; 1602 1596 }; ··· 1611 1605 error.CompressionError => conn.writeFramed(CLOSE_PROTOCOL_ERROR) catch {}, 1612 1606 else => {}, 1613 1607 } 1614 - log.debug("({}) invalid websocket packet: {}", .{ conn.address, err }); 1608 + log.debug("({f}) invalid websocket packet: {}", .{ conn.address, err }); 1615 1609 return false; 1616 1610 } orelse { 1617 1611 // everything is fine, we just need more data ··· 1621 1615 const message_type = message.type; 1622 1616 defer reader.done(message_type); 1623 1617 1624 - log.debug("({}) received {s} message", .{ hc.conn.address, @tagName(message_type) }); 1618 + log.debug("({f}) received {s} message", .{ hc.conn.address, @tagName(message_type) }); 1625 1619 switch (message_type) { 1626 1620 .text, .binary => { 1627 1621 const params = @typeInfo(@TypeOf(H.clientMessage)).@"fn".params; ··· 1997 1991 } 1998 1992 if (std.mem.eql(u8, data, "writer")) { 1999 1993 var wb = self.conn.writeBuffer(allocator, .text); 2000 - try std.fmt.format(wb.writer(), "{d}!!!", .{9000}); 2001 - return wb.flush(); 1994 + try wb.interface.print("{d}!!!", .{9000}); 1995 + return wb.send(); 2002 1996 } 2003 1997 if (std.mem.eql(u8, data, "ping")) { 2004 1998 var buf = [_]u8{ 'a', '-', 'p', 'i', 'n', 'g' };
+3 -3
src/server/thread_pool.zig
··· 195 195 tp.spawn(.{1}); 196 196 } 197 197 while (tp.empty() == false) { 198 - std.time.sleep(std.time.ns_per_ms); 198 + std.Thread.sleep(std.time.ns_per_ms); 199 199 } 200 200 tp.deinit(); 201 201 try t.expectEqual(50_000, testSum); ··· 209 209 tp.spawn(.{1}); 210 210 } 211 211 while (tp.empty() == false) { 212 - std.time.sleep(std.time.ns_per_ms); 212 + std.Thread.sleep(std.time.ns_per_ms); 213 213 } 214 214 tp.deinit(); 215 215 try t.expectEqual(50_000, testSum); ··· 220 220 std.debug.assert(buf.len == 512); 221 221 _ = @atomicRmw(u64, &testSum, .Add, c, .monotonic); 222 222 // let the threadpool queue get backed up 223 - std.time.sleep(std.time.ns_per_us * 100); 223 + std.Thread.sleep(std.time.ns_per_us * 100); 224 224 }
+4 -4
src/t.zig
··· 35 35 pub fn init() Writer { 36 36 return .{ 37 37 .pos = 0, 38 + .buf = .empty, 38 39 .random = getRandom(), 39 - .buf = std.ArrayList(u8).init(allocator), 40 40 }; 41 41 } 42 42 43 - pub fn deinit(self: *const Writer) void { 44 - self.buf.deinit(); 43 + pub fn deinit(self: *Writer) void { 44 + self.buf.deinit(allocator); 45 45 } 46 46 47 47 pub fn ping(self: *Writer) void { ··· 80 80 81 81 // 2 byte header + length_of_length + mask + payload_length 82 82 const needed = 2 + length_of_length + 4 + l; 83 - buf.ensureUnusedCapacity(needed) catch unreachable; 83 + buf.ensureUnusedCapacity(allocator, needed) catch unreachable; 84 84 85 85 if (fin) { 86 86 buf.appendAssumeCapacity(128 | op_code | reserved);
+4 -2
src/websocket.zig
··· 22 22 pub const Compression = struct { 23 23 write_threshold: ?usize = null, 24 24 retain_write_buffer: bool = true, 25 - client_no_context_takeover: bool = false, 26 - server_no_context_takeover: bool = false, 25 + // don't know how to support these with the Zig 0.15 changes. So, for now 26 + // we'll always require these to be true 27 + // client_no_context_takeover: bool = false, 28 + // server_no_context_takeover: bool = false, 27 29 }; 28 30 29 31 pub fn bufferProvider(allocator: std.mem.Allocator, config: buffer.Config) !buffer.Provider {
+5 -3
support/autobahn/client/build.zig
··· 6 6 7 7 const exe = b.addExecutable(.{ 8 8 .name = "autobahn_test_client", 9 - .root_source_file = b.path("main.zig"), 10 - .target = target, 11 - .optimize = optimize, 9 + .root_module = b.createModule(.{ 10 + .root_source_file = b.path("main.zig"), 11 + .target = target, 12 + .optimize = optimize, 13 + }), 12 14 }); 13 15 exe.root_module.addImport("websocket", b.dependency("websocket", .{}).module("websocket")); 14 16
+5 -4
support/autobahn/client/main.zig
··· 34 34 }; 35 35 36 36 // wait 5 seconds for autobanh server to be up 37 - std.time.sleep(std.time.ns_per_s * 5); 37 + std.Thread.sleep(std.time.ns_per_s * 5); 38 38 39 39 var buffer_provider = try websocket.bufferProvider(allocator, .{ .count = 10, .size = 32768, .max = 20_000_000 }); 40 40 defer buffer_provider.deinit(); ··· 92 92 .port = 9001, 93 93 .host = "localhost", 94 94 .buffer_provider = buffer_provider, 95 - .compression = .{ 96 - .write_threshold = 0, 97 - }, 95 + // zig 0.15 96 + // .compression = .{ 97 + // .write_threshold = 0, 98 + // }, 98 99 }); 99 100 errdefer client.deinit(); 100 101 try client.handshake(path, .{
+5 -3
support/autobahn/server/build.zig
··· 6 6 7 7 const exe = b.addExecutable(.{ 8 8 .name = "autobahn_test_server", 9 - .root_source_file = b.path("main.zig"), 10 - .target = target, 11 - .optimize = optimize, 9 + .root_module = b.createModule(.{ 10 + .root_source_file = b.path("main.zig"), 11 + .target = target, 12 + .optimize = optimize, 13 + }), 12 14 }); 13 15 14 16 const websocket = b.dependency("websocket", .{}).module("websocket");
+1 -2
support/autobahn/server/config.json
··· 2 2 "outdir": "/ab/reports/", 3 3 "options": {"failByDrop": false}, 4 4 "servers": [ 5 - {"agent": "non-blocking", "url": "ws://host.docker.internal:9224"}, 6 - {"agent": "non-blocking buffer pool", "url": "ws://host.docker.internal:9225"} 5 + {"agent": "non-blocking", "url": "ws://host.docker.internal:9224"} 7 6 ], 8 7 "cases": ["*"], 9 8 "exclude-cases": [],
+10 -8
support/autobahn/server/main.zig
··· 22 22 23 23 std.posix.sigaction(std.posix.SIG.TERM, &.{ 24 24 .handler = .{ .handler = shutdown }, 25 - .mask = std.posix.empty_sigset, 25 + .mask = std.posix.sigemptyset(), 26 26 .flags = 0, 27 27 }, null); 28 28 } ··· 53 53 .max_size = 1024, 54 54 .max_headers = 10, 55 55 }, 56 - .compression = .{ 57 - .write_threshold = 0, 58 - }, 56 + // zig 0.15 57 + // .compression = .{ 58 + // .write_threshold = 0, 59 + // }, 59 60 }); 60 61 return try nonblocking_server.listenInNewThread({}); 61 62 } ··· 76 77 .max_size = 1024, 77 78 .max_headers = 10, 78 79 }, 79 - .compression = .{ 80 - .write_threshold = 0, 81 - }, 80 + // zig 0.15 81 + // .compression = .{ 82 + // .write_threshold = 0, 83 + // }, 82 84 }); 83 85 return try nonblocking_bp_server.listenInNewThread({}); 84 86 } ··· 107 109 } 108 110 }; 109 111 110 - fn shutdown(_: c_int) callconv(.C) void { 112 + fn shutdown(_: c_int) callconv(.c) void { 111 113 nonblocking_server.stop(); 112 114 nonblocking_bp_server.stop(); 113 115 }
+28 -42
test_runner.zig
··· 1 1 // in your build.zig, you can specify a custom test runner: 2 2 // const tests = b.addTest(.{ 3 - // .target = target, 4 - // .optimize = optimize, 5 - // .test_runner = .{ .path = b.path("test_runner.zig"), .mode = .simple }, // add this line 6 - // .root_source_file = b.path("src/main.zig"), 3 + // .root_module = $MODULE_BEING_TESTED, 4 + // .test_runner = .{ .path = b.path("test_runner.zig"), .mode = .simple }, 7 5 // }); 8 6 9 7 pub const std_options = std.Options{ .log_scope_levels = &[_]std.log.ScopeLevel{ ··· 37 35 var skip: usize = 0; 38 36 var leak: usize = 0; 39 37 40 - const printer = Printer.init(); 41 - printer.fmt("\r\x1b[0K", .{}); // beginning of line and clear to end of line 38 + Printer.fmt("\r\x1b[0K", .{}); // beginning of line and clear to end of line 42 39 43 40 for (builtin.test_functions) |t| { 44 41 if (isSetup(t)) { 45 42 t.func() catch |err| { 46 - printer.status(.fail, "\nsetup \"{s}\" failed: {}\n", .{ t.name, err }); 43 + Printer.status(.fail, "\nsetup \"{s}\" failed: {}\n", .{ t.name, err }); 47 44 return err; 48 45 }; 49 46 } ··· 85 82 86 83 if (std.testing.allocator_instance.deinit() == .leak) { 87 84 leak += 1; 88 - printer.status(.fail, "\n{s}\n\"{s}\" - Memory Leak\n{s}\n", .{ BORDER, friendly_name, BORDER }); 85 + Printer.status(.fail, "\n{s}\n\"{s}\" - Memory Leak\n{s}\n", .{ BORDER, friendly_name, BORDER }); 89 86 } 90 87 91 88 if (result) |_| { ··· 98 95 else => { 99 96 status = .fail; 100 97 fail += 1; 101 - printer.status(.fail, "\n{s}\n\"{s}\" - {s}\n{s}\n", .{ BORDER, friendly_name, @errorName(err), BORDER }); 98 + Printer.status(.fail, "\n{s}\n\"{s}\" - {s}\n{s}\n", .{ BORDER, friendly_name, @errorName(err), BORDER }); 102 99 if (@errorReturnTrace()) |trace| { 103 100 std.debug.dumpStackTrace(trace.*); 104 101 } ··· 110 107 111 108 if (env.verbose) { 112 109 const ms = @as(f64, @floatFromInt(ns_taken)) / 1_000_000.0; 113 - printer.status(status, "{s} ({d:.2}ms)\n", .{ friendly_name, ms }); 110 + Printer.status(status, "{s} ({d:.2}ms)\n", .{ friendly_name, ms }); 114 111 } else { 115 - printer.status(status, ".", .{}); 112 + Printer.status(status, ".", .{}); 116 113 } 117 114 } 118 115 119 116 for (builtin.test_functions) |t| { 120 117 if (isTeardown(t)) { 121 118 t.func() catch |err| { 122 - printer.status(.fail, "\nteardown \"{s}\" failed: {}\n", .{ t.name, err }); 119 + Printer.status(.fail, "\nteardown \"{s}\" failed: {}\n", .{ t.name, err }); 123 120 return err; 124 121 }; 125 122 } ··· 127 124 128 125 const total_tests = pass + fail; 129 126 const status = if (fail == 0) Status.pass else Status.fail; 130 - printer.status(status, "\n{d} of {d} test{s} passed\n", .{ pass, total_tests, if (total_tests != 1) "s" else "" }); 127 + Printer.status(status, "\n{d} of {d} test{s} passed\n", .{ pass, total_tests, if (total_tests != 1) "s" else "" }); 131 128 if (skip > 0) { 132 - printer.status(.skip, "{d} test{s} skipped\n", .{ skip, if (skip != 1) "s" else "" }); 129 + Printer.status(.skip, "{d} test{s} skipped\n", .{ skip, if (skip != 1) "s" else "" }); 133 130 } 134 131 if (leak > 0) { 135 - printer.status(.fail, "{d} test{s} leaked\n", .{ leak, if (leak != 1) "s" else "" }); 132 + Printer.status(.fail, "{d} test{s} leaked\n", .{ leak, if (leak != 1) "s" else "" }); 136 133 } 137 - printer.fmt("\n", .{}); 138 - try slowest.display(printer); 139 - printer.fmt("\n", .{}); 134 + Printer.fmt("\n", .{}); 135 + try slowest.display(); 136 + Printer.fmt("\n", .{}); 140 137 std.posix.exit(if (fail == 0) 0 else 1); 141 138 } 142 139 143 140 const Printer = struct { 144 - out: std.fs.File.Writer, 145 - 146 - fn init() Printer { 147 - return .{ 148 - .out = std.io.getStdErr().writer(), 149 - }; 141 + fn fmt(comptime format: []const u8, args: anytype) void { 142 + std.debug.print(format, args); 150 143 } 151 144 152 - fn fmt(self: Printer, comptime format: []const u8, args: anytype) void { 153 - std.fmt.format(self.out, format, args) catch unreachable; 154 - } 155 - 156 - fn status(self: Printer, s: Status, comptime format: []const u8, args: anytype) void { 157 - const color = switch (s) { 158 - .pass => "\x1b[32m", 159 - .fail => "\x1b[31m", 160 - .skip => "\x1b[33m", 161 - else => "", 162 - }; 163 - const out = self.out; 164 - out.writeAll(color) catch @panic("writeAll failed?!"); 165 - std.fmt.format(out, format, args) catch @panic("std.fmt.format failed?!"); 166 - self.fmt("\x1b[0m", .{}); 145 + fn status(s: Status, comptime format: []const u8, args: anytype) void { 146 + switch (s) { 147 + .pass => std.debug.print("\x1b[32m", .{}), 148 + .fail => std.debug.print("\x1b[31m", .{}), 149 + .skip => std.debug.print("\x1b[33m", .{}), 150 + else => {}, 151 + } 152 + std.debug.print(format ++ "\x1b[0m", args); 167 153 } 168 154 }; 169 155 ··· 233 219 return ns; 234 220 } 235 221 236 - fn display(self: *SlowTracker, printer: Printer) !void { 222 + fn display(self: *SlowTracker) !void { 237 223 var slowest = self.slowest; 238 224 const count = slowest.count(); 239 - printer.fmt("Slowest {d} test{s}: \n", .{ count, if (count != 1) "s" else "" }); 225 + Printer.fmt("Slowest {d} test{s}: \n", .{ count, if (count != 1) "s" else "" }); 240 226 while (slowest.removeMinOrNull()) |info| { 241 227 const ms = @as(f64, @floatFromInt(info.ns)) / 1_000_000.0; 242 - printer.fmt(" {d:.2}ms\t{s}\n", .{ ms, info.name }); 228 + Printer.fmt(" {d:.2}ms\t{s}\n", .{ ms, info.name }); 243 229 } 244 230 } 245 231