websocket
0

Configure Feed

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

websocket.zig / src / client / client.zig
47 kB 1296 lines
1const std = @import("std"); 2const proto = @import("../proto.zig"); 3const buffer = @import("../buffer.zig"); 4 5const ascii = std.ascii; 6const Io = std.Io; 7const net = Io.net; 8const posix = std.posix; 9const tls = std.crypto.tls; 10const log = std.log.scoped(.websocket); 11 12const Reader = proto.Reader; 13const Allocator = std.mem.Allocator; 14const Bundle = std.crypto.Certificate.Bundle; 15const CompressionOpts = @import("../websocket.zig").Compression; 16const ServerHandshake = @import("../server/handshake.zig").Handshake; 17 18fn milliTimestamp(io: Io) i64 { 19 const ts = Io.Timestamp.now(io, .real); 20 return @intCast(@divTrunc(ts.nanoseconds, std.time.ns_per_ms)); 21} 22 23fn ReadLoopHandler(comptime T: type) type { 24 const info = @typeInfo(T); 25 26 switch (info) { 27 .@"struct" => |struct_info| { 28 if (struct_info.is_tuple) 29 @compileError("readLoop: handler does not support tuples."); 30 31 return T; 32 }, 33 .pointer => |ptr_info| { 34 switch (ptr_info.size) { 35 .one => return ReadLoopHandler(ptr_info.child), 36 else => @compileError("readLoop: handler does not support Slice, C and Many pointers."), 37 } 38 }, 39 else => @compileError("readLoop: expected handler to be a struct or pointer to a struct but found '" ++ @tagName(info) ++ "'"), 40 } 41} 42 43pub const Client = struct { 44 io: Io, 45 stream: Stream, 46 _reader: Reader, 47 _closed: bool, 48 _compression_opts: ?CompressionOpts, 49 _compression: ?Client.Compression = null, 50 51 // Serializes writes from concurrent tasks (ping loop, auto-pong, close). 52 // Matches server-side Conn.lock pattern. 53 _write_lock: Io.Mutex = .init, 54 55 // When creating a client, we can either be given a BufferProvider or create 56 // one ourselves. If we create it ourselves (in init), we "own" it and must 57 // free it on deinit. (The reference to the buffer provider is already in the 58 // reader, no need to hold another reference in the client). 59 _own_bp: bool, 60 61 // For advanced cases, a custom masking function can be provided. Masking 62 // is a security feature that only really makes sense in the browser. If you 63 // aren't running websockets in the browser AND you control both the client 64 // and the server, you could get a performance boost by not masking. 65 _mask_fn: *const fn (Io) [4]u8, 66 67 pub const Config = struct { 68 port: u16, 69 host: []const u8, 70 tls: bool = false, 71 max_size: usize = 65536, 72 buffer_size: usize = 4096, 73 ca_bundle: ?Bundle = null, 74 mask_fn: ?*const fn (Io) [4]u8 = null, 75 buffer_provider: ?*buffer.Provider = null, 76 compression: ?CompressionOpts = null, 77 }; 78 79 pub const HandshakeOpts = struct { 80 timeout_ms: u32 = 10000, 81 headers: ?[]const u8 = null, 82 }; 83 84 const Compression = struct { 85 allocator: Allocator, 86 retain_writer: bool, 87 write_treshold: usize, 88 writer: std.Io.Writer.Allocating, 89 }; 90 91 pub fn init(io: Io, allocator: Allocator, config: Config) !Client { 92 if (config.compression != null) { 93 log.err("Compression is disabled as part of the 0.15 upgrade. I do hope to re-enable it soon.", .{}); 94 return error.InvalidConfiguraion; 95 } 96 97 // 0.16: networking via Io.net.HostName 98 const host_name = try net.HostName.init(config.host); 99 // 0.16: connect requires mode option (stream vs datagram) 100 const net_stream = try host_name.connect(io, config.port, .{ .mode = .stream }); 101 102 var tls_client: ?*TLSClient = null; 103 if (config.tls) { 104 tls_client = try TLSClient.init(allocator, io, net_stream, &config); 105 } 106 const stream = Stream.init(io, net_stream, tls_client); 107 108 var own_bp = false; 109 var buffer_provider: *buffer.Provider = undefined; 110 111 // If a buffer_provider is provided, we'll use that. 112 // If it isn't, we need to create one which also means we now "own" it 113 // and we're responsible for cleaning it up 114 if (config.buffer_provider) |shared_bp| { 115 buffer_provider = shared_bp; 116 } else { 117 own_bp = true; 118 buffer_provider = try allocator.create(buffer.Provider); 119 errdefer allocator.destroy(buffer_provider); 120 buffer_provider.* = try buffer.Provider.init(allocator, .{ 121 .size = 0, 122 .count = 0, 123 .max = config.max_size, 124 }); 125 } 126 127 errdefer if (own_bp) { 128 buffer_provider.deinit(); 129 allocator.destroy(buffer_provider); 130 }; 131 132 const reader_buf = try buffer_provider.allocator.alloc(u8, config.buffer_size); 133 errdefer buffer_provider.allocator.free(reader_buf); 134 135 return .{ 136 .io = io, 137 .stream = stream, 138 ._closed = false, 139 ._own_bp = own_bp, 140 ._mask_fn = config.mask_fn orelse generateMask, 141 ._compression_opts = null, //TODO: ZIG 0.15 142 ._reader = Reader.init(reader_buf, buffer_provider, null), 143 }; 144 } 145 146 // 0.16: Alternative init that accepts a pre-existing stream. 147 // Supports TLS if config.tls is set (uses config.host for SNI). 148 pub fn initWithStream(io: Io, allocator: Allocator, net_stream: net.Stream, config: Config) !Client { 149 var tls_client: ?*TLSClient = null; 150 if (config.tls) { 151 tls_client = try TLSClient.init(allocator, io, net_stream, &config); 152 } 153 const stream = Stream.init(io, net_stream, tls_client); 154 155 var own_bp = false; 156 var buffer_provider: *buffer.Provider = undefined; 157 158 if (config.buffer_provider) |shared_bp| { 159 buffer_provider = shared_bp; 160 } else { 161 own_bp = true; 162 buffer_provider = try allocator.create(buffer.Provider); 163 errdefer allocator.destroy(buffer_provider); 164 buffer_provider.* = try buffer.Provider.init(allocator, .{ 165 .size = 0, 166 .count = 0, 167 .max = config.max_size, 168 }); 169 } 170 171 errdefer if (own_bp) { 172 buffer_provider.deinit(); 173 allocator.destroy(buffer_provider); 174 }; 175 176 const reader_buf = try buffer_provider.allocator.alloc(u8, config.buffer_size); 177 errdefer buffer_provider.allocator.free(reader_buf); 178 179 return .{ 180 .io = io, 181 .stream = stream, 182 ._closed = false, 183 ._own_bp = own_bp, 184 ._mask_fn = config.mask_fn orelse generateMask, 185 ._compression_opts = null, 186 ._reader = Reader.init(reader_buf, buffer_provider, null), 187 }; 188 } 189 190 pub fn deinit(self: *Client) void { 191 self.closeStream(); 192 193 const larger_buffer_provider = self._reader.large_buffer_provider; 194 const allocator = larger_buffer_provider.allocator; 195 allocator.free(self._reader.static); 196 197 self._reader.deinit(); 198 199 if (self._own_bp) { 200 larger_buffer_provider.deinit(); 201 allocator.destroy(larger_buffer_provider); 202 } 203 } 204 205 pub fn handshake(self: *Client, path: []const u8, opts: HandshakeOpts) !void { 206 const stream = &self.stream; 207 errdefer self.closeStream(); 208 209 // we've already setup our reader, and the reader has a static buffer 210 // we might as well use it! 211 const buf = self._reader.static; 212 const key = blk: { 213 const bin_key = generateKey(self.io); 214 var encoded_key: [24]u8 = undefined; 215 break :blk std.base64.standard.Encoder.encode(&encoded_key, &bin_key); 216 }; 217 218 try sendHandshake(path, key, buf, &opts, self._compression_opts != null, stream); 219 220 const res = try HandShakeReply.read(buf, key, &opts, self._compression_opts != null, stream); 221 errdefer self.close(.{ .code = 1001 }) catch unreachable; 222 223 // Set up compression with agreed-on parameters 224 if (res.compression) { 225 try self.setupCompression(); 226 } 227 228 // We might have read more than handshake response. If so, readHandshakeReply 229 // has positioned the extra data at the start of the buffer, but we need 230 // to set the length. 231 self._reader.pos = res.over_read; 232 } 233 234 fn setupCompression(self: *Client) !void { 235 std.debug.assert(self._compression_opts != null); 236 self._reader.allow_compressed = true; 237 238 const allocator = self._reader.large_buffer_provider.allocator; 239 const config = self._compression_opts.?; 240 self._compression = .{ 241 .allocator = allocator, 242 .write_treshold = config.write_threshold.?, 243 .retain_writer = config.retain_write_buffer, 244 .writer = std.Io.Writer.Allocating.init(allocator), 245 }; 246 } 247 248 pub fn readLoop(self: *Client, handler: anytype) !void { 249 const Handler = ReadLoopHandler(@TypeOf(handler)); 250 var reader = &self._reader; 251 252 defer if (comptime std.meta.hasFn(Handler, "close")) { 253 handler.close(); 254 }; 255 256 // block until we have data 257 try self.readTimeout(0); 258 259 while (true) { 260 const message = self.read() catch |err| switch (err) { 261 error.Closed => return, 262 else => return err, 263 } orelse unreachable; 264 265 const message_type = message.type; 266 defer reader.done(message_type); 267 268 switch (message_type) { 269 .text, .binary => { 270 switch (comptime @typeInfo(@TypeOf(Handler.serverMessage)).@"fn".param_types.len) { 271 2 => try handler.serverMessage(message.data), 272 3 => try handler.serverMessage(message.data, if (message_type == .text) .text else .binary), 273 else => @compileError(@typeName(Handler) ++ ".serverMessage must accept 2 or 3 parameters"), 274 } 275 }, 276 .ping => if (comptime std.meta.hasFn(Handler, "serverPing")) { 277 try handler.serverPing(message.data); 278 } else { 279 // @constCast is safe because we know message.data points to 280 // reader.buffer.buf, which we own and which can be mutated 281 try self.writeFrame(.pong, @constCast(message.data)); 282 }, 283 .close => { 284 if (comptime std.meta.hasFn(Handler, "serverClose")) { 285 try handler.serverClose(message.data); 286 } else { 287 self.close(.{}) catch unreachable; 288 } 289 return; 290 }, 291 .pong => if (comptime std.meta.hasFn(Handler, "serverPong")) { 292 try handler.serverPong(message.data); 293 }, 294 } 295 } 296 } 297 298 pub const HeartbeatConfig = struct { 299 /// ping interval in milliseconds. readTimeout is set to this value. 300 /// when no data arrives within the interval, a ping is sent. 301 interval_ms: u32 = 30_000, 302 /// close connection after this many consecutive intervals with no data or pong. 303 max_failures: u32 = 4, 304 }; 305 306 pub fn readLoopWithHeartbeat(self: *Client, handler: anytype, heartbeat: HeartbeatConfig) !void { 307 const Handler = ReadLoopHandler(@TypeOf(handler)); 308 var reader = &self._reader; 309 310 defer if (comptime std.meta.hasFn(Handler, "close")) { 311 handler.close(); 312 }; 313 314 try self.readTimeout(heartbeat.interval_ms); 315 316 var pending_pings: u32 = 0; 317 318 while (true) { 319 const message = self.read() catch |err| switch (err) { 320 error.Closed => return, 321 else => return err, 322 } orelse { 323 // timeout — no data in interval 324 pending_pings += 1; 325 if (pending_pings >= heartbeat.max_failures) { 326 self.close(.{}) catch {}; 327 return error.Closed; 328 } 329 self.writePing(&.{}) catch { 330 self.close(.{}) catch {}; 331 return error.Closed; 332 }; 333 continue; 334 }; 335 336 // any received frame proves liveness 337 pending_pings = 0; 338 339 const message_type = message.type; 340 defer reader.done(message_type); 341 342 switch (message_type) { 343 .text, .binary => { 344 switch (comptime @typeInfo(@TypeOf(Handler.serverMessage)).@"fn".param_types.len) { 345 2 => try handler.serverMessage(message.data), 346 3 => try handler.serverMessage(message.data, if (message_type == .text) .text else .binary), 347 else => @compileError(@typeName(Handler) ++ ".serverMessage must accept 2 or 3 parameters"), 348 } 349 }, 350 .ping => if (comptime std.meta.hasFn(Handler, "serverPing")) { 351 try handler.serverPing(message.data); 352 } else { 353 try self.writeFrame(.pong, @constCast(message.data)); 354 }, 355 .close => { 356 if (comptime std.meta.hasFn(Handler, "serverClose")) { 357 try handler.serverClose(message.data); 358 } else { 359 self.close(.{}) catch unreachable; 360 } 361 return; 362 }, 363 .pong => if (comptime std.meta.hasFn(Handler, "serverPong")) { 364 try handler.serverPong(message.data); 365 }, 366 } 367 } 368 } 369 370 pub fn read(self: *Client) !?proto.Message { 371 var reader = &self._reader; 372 const stream = &self.stream; 373 374 while (true) { 375 // try to read a message from our buffer first, before trying to 376 // get more data from the socket. 377 const has_more, const message = reader.read() catch |err| { 378 self.close(.{ .code = 1002 }) catch unreachable; 379 return err; 380 } orelse { 381 // 0.16: Io vtable error set changed 382 reader.fill(stream) catch |err| { 383 // Check for timeout/would-block type errors 384 if (err == error.Canceled or err == error.WouldBlock) return null; 385 // Check for connection closed errors 386 if (err == error.Closed or err == error.ConnectionResetByPeer or err == error.NotOpenForReading) { 387 @atomicStore(bool, &self._closed, true, .monotonic); 388 return error.Closed; 389 } 390 self.close(.{ .code = 1002 }) catch unreachable; 391 return err; 392 }; 393 continue; 394 }; 395 396 _ = has_more; 397 return message; 398 } 399 } 400 401 pub fn done(self: *Client, message: proto.Message) void { 402 self._reader.done(message.type); 403 } 404 405 pub fn readLoopInNewThread(self: *Client, h: anytype) !std.Thread { 406 return std.Thread.spawn(.{}, readLoopOwnedThread, .{ self, h }); 407 } 408 409 fn readLoopOwnedThread(self: *Client, h: anytype) void { 410 self.readLoop(h) catch {}; 411 } 412 413 pub fn writeTimeout(self: *const Client, ms: u32) !void { 414 return self.stream.writeTimeout(ms); 415 } 416 417 pub fn readTimeout(self: *const Client, ms: u32) !void { 418 return self.stream.readTimeout(ms); 419 } 420 421 pub fn write(self: *Client, data: []u8) !void { 422 return self.writeFrame(.text, data); 423 } 424 425 pub fn writeText(self: *Client, data: []u8) !void { 426 return self.writeFrame(.text, data); 427 } 428 429 pub fn writeBin(self: *Client, data: []u8) !void { 430 return self.writeFrame(.binary, data); 431 } 432 433 pub fn writePing(self: *Client, data: []u8) !void { 434 return self.writeFrame(.ping, data); 435 } 436 437 pub fn writePong(self: *Client, data: []u8) !void { 438 return self.writeFrame(.pong, data); 439 } 440 441 const CloseOpts = struct { 442 code: ?u16 = null, 443 reason: []const u8 = "", 444 }; 445 446 pub fn close(self: *Client, opts: CloseOpts) !void { 447 if (@atomicRmw(bool, &self._closed, .Xchg, true, .monotonic) == true) { 448 // already closed 449 return; 450 } 451 452 defer self.stream.close(); 453 454 const code = opts.code orelse { 455 self.writeFrame(.close, "") catch {}; 456 return; 457 }; 458 459 const reason = opts.reason; 460 if (reason.len > 123) { 461 return error.ReasonTooLong; 462 } 463 464 var buf: [125]u8 = undefined; 465 buf[0] = @intCast((code >> 8) & 0xFF); 466 buf[1] = @intCast(code & 0xFF); 467 468 const end = 2 + reason.len; 469 @memcpy(buf[2..end], reason); 470 self.writeFrame(.close, buf[0..end]) catch {}; 471 } 472 473 pub fn writeFrame(self: *Client, op_code: proto.OpCode, data: []u8) !void { 474 const payload = data; 475 const compressed = false; 476 477 // maximum possible prefix length. op_code + length_type + 8byte length + 4 byte mask 478 var buf: [14]u8 = undefined; 479 const header = proto.writeFrameHeader(&buf, op_code, payload.len, compressed); 480 481 const header_len = header.len; 482 const header_end = header.len + 4; // for the mask 483 484 buf[1] |= 128; // indicate that the payload is masked 485 486 const mask = self._mask_fn(self.io); 487 @memcpy(buf[header_len..header_end], &mask); 488 489 if (payload.len > 0) { 490 proto.mask(&mask, payload); 491 } 492 493 // Serialize writes — concurrent ping/pong/close must not interleave frames. 494 self._write_lock.lockUncancelable(self.io); 495 defer self._write_lock.unlock(self.io); 496 497 try self.stream.writeAll(buf[0..header_end]); 498 if (payload.len > 0) { 499 try self.stream.writeAll(payload); 500 } 501 } 502 503 pub fn isClosed(self: *const Client) bool { 504 return @atomicLoad(bool, &self._closed, .monotonic); 505 } 506 507 fn closeStream(self: *Client) void { 508 if (@atomicRmw(bool, &self._closed, .Xchg, true, .monotonic) == false) { 509 self.stream.close(); 510 } 511 } 512}; 513 514// wraps a net.Stream and optional a tls.Client 515pub const Stream = struct { 516 io: Io, 517 stream: net.Stream, 518 tls_client: ?*TLSClient = null, 519 520 pub fn init(io: Io, stream: net.Stream, tls_client: ?*TLSClient) Stream { 521 return .{ 522 .io = io, 523 .stream = stream, 524 .tls_client = tls_client, 525 }; 526 } 527 528 pub fn close(self: *Stream) void { 529 if (self.tls_client) |tls_client| { 530 self.stream.shutdown(self.io, .both) catch {}; 531 tls_client.deinit(); 532 } 533 self.stream.close(self.io); 534 } 535 536 pub fn read(self: *Stream, buf: []u8) !usize { 537 if (self.tls_client) |tls_client| { 538 var w: std.Io.Writer = .fixed(buf); 539 while (true) { 540 const n = try tls_client.client.reader.stream(&w, .limited(buf.len)); 541 if (n != 0) { 542 return n; 543 } 544 } 545 } 546 return posix.read(self.stream.socket.handle, buf) catch |err| { 547 return switch (err) { 548 error.ConnectionResetByPeer => error.ConnectionResetByPeer, 549 error.WouldBlock => error.WouldBlock, 550 else => error.Unexpected, 551 }; 552 }; 553 } 554 555 pub fn writeAll(self: *Stream, data: []const u8) !void { 556 if (self.tls_client) |tls_client| { 557 try tls_client.client.writer.writeAll(data); 558 try tls_client.client.writer.flush(); 559 try tls_client.stream_writer.interface.flush(); 560 return; 561 } 562 var remaining = data; 563 while (remaining.len > 0) { 564 // netWrite: header is sent first, data array's last element is the splat pattern. 565 // Pass remaining as header, empty pattern with splat=0. 566 const empty = [_][]const u8{""}; 567 const n = self.io.vtable.netWrite(self.io.userdata, self.stream.socket.handle, remaining, &empty, 0) catch |err| { 568 return switch (err) { 569 error.ConnectionResetByPeer => error.ConnectionResetByPeer, 570 else => error.Unexpected, 571 }; 572 }; 573 if (n == 0) return error.Unexpected; 574 remaining = remaining[n..]; 575 } 576 } 577 578 const zero_timeout = std.mem.toBytes(posix.timeval{ .sec = 0, .usec = 0 }); 579 pub fn writeTimeout(self: *const Stream, ms: u32) !void { 580 return self.setTimeout(posix.SO.SNDTIMEO, ms); 581 } 582 583 pub fn readTimeout(self: *const Stream, ms: u32) !void { 584 return self.setTimeout(posix.SO.RCVTIMEO, ms); 585 } 586 587 fn setTimeout(self: *const Stream, opt_name: u32, ms: u32) !void { 588 if (ms == 0) { 589 return self.setsockopt(opt_name, &zero_timeout); 590 } 591 592 const timeout = std.mem.toBytes(posix.timeval{ 593 .sec = @intCast(@divTrunc(ms, 1000)), 594 .usec = @intCast(@mod(ms, 1000) * 1000), 595 }); 596 return self.setsockopt(opt_name, &timeout); 597 } 598 599 pub fn setsockopt(self: *const Stream, opt_name: u32, value: []const u8) !void { 600 return setConnSockOpt(self.stream.socket.handle, posix.SOL.SOCKET, opt_name, value); 601 } 602}; 603 604/// setsockopt for a *connection* socket. 605/// 606/// std.posix.setsockopt maps BADF/NOTSOCK/INVAL/FAULT to `unreachable` 607/// ("always a race condition"). On a listening socket that holds. On a 608/// connection socket it does not: the peer can reset, or another thread can 609/// close the fd, between connect and the timeout call. `unreachable` cannot be 610/// caught, so that race aborts the whole process instead of surfacing as an 611/// error the caller's reconnect path would handle — observed as a panic inside 612/// the client handshake when an upstream dropped the connection. 613/// 614/// FAULT is deliberately NOT swallowed: a bad option pointer is our own bug, 615/// not a peer's doing, and hiding it would trade one silent failure for 616/// another. 617/// 618/// The timeouts this backs are advisory: if the socket really is gone, the 619/// following read or write reports it properly. So swallow exactly the arms the 620/// stdlib calls impossible and keep every other error meaningful. 621/// 622/// The server side took the same approach in `setSockOptBestEffort` 623/// (src/server/server.zig). 624fn setConnSockOpt(fd: posix.socket_t, level: i32, optname: u32, opt: []const u8) !void { 625 switch (posix.errno(posix.system.setsockopt(fd, level, optname, opt.ptr, @intCast(opt.len)))) { 626 .SUCCESS => {}, 627 // The socket died under us. Not our failure to report. 628 .BADF, .NOTSOCK, .INVAL => {}, 629 .DOM => return error.TimeoutTooBig, 630 .ISCONN => return error.AlreadyConnected, 631 .NOPROTOOPT => return error.InvalidProtocolOption, 632 .NOMEM, .NOBUFS => return error.SystemResources, 633 .PERM => return error.PermissionDenied, 634 .NODEV => return error.NoDevice, 635 .OPNOTSUPP => return error.OperationUnsupported, 636 else => |err| return posix.unexpectedErrno(err), 637 } 638} 639 640const TLSClient = struct { 641 client: tls.Client, 642 stream: net.Stream, 643 stream_writer: net.Stream.Writer, 644 stream_reader: net.Stream.Reader, 645 arena: std.heap.ArenaAllocator, 646 647 fn init(allocator: Allocator, io: Io, stream: net.Stream, config: *const Client.Config) !*TLSClient { 648 var arena = std.heap.ArenaAllocator.init(allocator); 649 errdefer arena.deinit(); 650 651 const aa = arena.allocator(); 652 653 const bundle_ptr = try aa.create(Bundle); 654 if (config.ca_bundle) |ca| { 655 bundle_ptr.* = ca; 656 } else { 657 bundle_ptr.* = .{ .map = .empty, .bytes = .empty }; 658 try bundle_ptr.rescan(aa, io, Io.Timestamp.zero); 659 } 660 661 const rwlock = try aa.create(std.Io.RwLock); 662 rwlock.* = std.Io.RwLock.init; 663 664 // The TLS input and output have to be max_ciphertext_record_len each. 665 // It isn't clear to me how big the un-encrypted reader and writer 666 // need to be. I would think 0, but that will fail an assertion. I 667 // don't think that it's right that we need 4 buffers, but apparently 668 // we do. Until i figure this out, using 4 x max_ciphertext_record_len 669 // seems like the only safe choice. 670 const buf_len = std.crypto.tls.max_ciphertext_record_len; 671 var buf = try aa.alloc(u8, buf_len * 4); 672 673 const self = try aa.create(TLSClient); 674 self.* = .{ 675 .stream = stream, 676 .arena = arena, 677 .client = undefined, 678 .stream_writer = stream.writer(io, buf.ptr[0..buf_len][0..buf_len]), 679 .stream_reader = stream.reader(io, buf.ptr[buf_len .. 2 * buf_len][0..buf_len]), 680 }; 681 682 var entropy: [tls.Client.Options.entropy_len]u8 = undefined; 683 io.random(&entropy); 684 self.client = try tls.Client.init( 685 &self.stream_reader.interface, 686 &self.stream_writer.interface, 687 .{ 688 .ca = .{ .bundle = .{ .gpa = aa, .io = io, .lock = rwlock, .bundle = bundle_ptr } }, 689 .host = .{ .explicit = config.host }, 690 .read_buffer = buf.ptr[2 * buf_len .. 3 * buf_len][0..buf_len], 691 .write_buffer = buf.ptr[3 * buf_len .. 4 * buf_len][0..buf_len], 692 .entropy = &entropy, 693 .realtime_now = Io.Timestamp.now(io, .real), 694 }, 695 ); 696 697 return self; 698 } 699 700 fn deinit(self: *TLSClient) void { 701 _ = self.client.end() catch {}; 702 self.arena.deinit(); 703 } 704}; 705 706fn generateKey(io: Io) [16]u8 { 707 if (comptime @import("builtin").is_test) { 708 return [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; 709 } 710 // 0.16: io.random() fills a buffer 711 var key: [16]u8 = undefined; 712 io.random(&key); 713 return key; 714} 715 716fn generateMask(io: Io) [4]u8 { 717 // 0.16: io.random() fills a buffer 718 var mask: [4]u8 = undefined; 719 io.random(&mask); 720 return mask; 721} 722 723fn sendHandshake(path: []const u8, key: []const u8, buf: []u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !void { 724 @memcpy(buf[0..4], "GET "); 725 var pos: usize = 4; 726 var end = pos + path.len; 727 728 { 729 @memcpy(buf[pos..end], path); 730 pos = end; 731 } 732 733 { 734 const headers = " HTTP/1.1\r\ncontent-length: 0\r\nupgrade: websocket\r\nsec-websocket-version: 13\r\nconnection: upgrade\r\nsec-websocket-key: "; 735 end = pos + headers.len; 736 @memcpy(buf[pos..end], headers); 737 738 pos = end; 739 end = pos + key.len; 740 @memcpy(buf[pos..end], key); 741 } 742 743 if (compression) { 744 // NOTE: client_max_window_bits is unsupported 745 const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover"; 746 pos = end; 747 end = pos + permessage_deflate.len; 748 @memcpy(buf[pos..end], permessage_deflate); 749 } 750 751 { 752 pos = end; 753 end = pos + 2; 754 @memcpy(buf[pos..end], "\r\n"); 755 pos = end; 756 } 757 758 if (opts.headers) |extra_headers| { 759 end = pos + extra_headers.len; 760 @memcpy(buf[pos..end], extra_headers); 761 pos = end; 762 if (!std.mem.endsWith(u8, extra_headers, "\r\n")) { 763 buf[pos] = '\r'; 764 buf[pos + 1] = '\n'; 765 pos += 2; 766 } 767 } 768 buf[pos] = '\r'; 769 buf[pos + 1] = '\n'; 770 771 try stream.writeTimeout(opts.timeout_ms); 772 try stream.writeAll(buf[0 .. pos + 2]); 773 try stream.writeTimeout(0); 774} 775 776const HandShakeReply = struct { 777 compression: bool, 778 over_read: usize, 779 780 fn read(buf: []u8, key: []const u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !HandShakeReply { 781 const timeout_ms = opts.timeout_ms; 782 const deadline = milliTimestamp(stream.io) + timeout_ms; 783 try stream.readTimeout(timeout_ms); 784 785 var pos: usize = 0; 786 var line_start: usize = 0; 787 var complete_response: u8 = 0; 788 var server_compression: bool = false; 789 790 while (true) { 791 // 0.16: using libc recv, WouldBlock indicates timeout 792 const n = stream.read(buf[pos..]) catch |err| switch (err) { 793 error.WouldBlock => return error.Timeout, 794 else => return err, 795 }; 796 if (n == 0) { 797 return error.ConnectionClosed; 798 } 799 800 pos += n; 801 while (std.mem.indexOfScalar(u8, buf[line_start..pos], '\r')) |relative_end| { 802 if (relative_end == 0) { 803 if (complete_response != 15) { 804 return error.InvalidHandshakeResponse; 805 } 806 // TCP can split the terminating CRLF — if the trailing \n 807 // hasn't arrived yet, pos == line_start + 1 and the over_read 808 // subtraction below would underflow. break for more data. 809 if (line_start + 2 > pos) break; 810 const over_read = pos - (line_start + 2); 811 std.mem.copyForwards(u8, buf[0..over_read], buf[line_start + 2 .. pos]); 812 try stream.readTimeout(0); 813 return .{ 814 .over_read = over_read, 815 .compression = server_compression, 816 }; 817 } 818 819 const line_end = line_start + relative_end; 820 const line = buf[line_start..line_end]; 821 822 // the next line starts where this line ends, skip over the \r\n. 823 // TCP can split mid-CRLF — if \n hasn't arrived yet, break to 824 // the outer read loop for more data. 825 line_start = line_end + 2; 826 if (line_start > pos) break; 827 828 if (complete_response == 0) { 829 if (!ascii.startsWithIgnoreCase(line, "HTTP/1.1 101 ")) { 830 return error.InvalidHandshakeResponse; 831 } 832 complete_response |= 1; 833 continue; 834 } 835 836 for (line, 0..) |b, i| { 837 // find the colon and lowercase the header while we're iterating 838 if ('A' <= b and b <= 'Z') { 839 line[i] = b + 32; 840 continue; 841 } 842 843 if (b != ':') { 844 continue; 845 } 846 847 switch (i) { 848 7 => if (std.mem.eql(u8, line[0..i], "upgrade")) { 849 if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "websocket")) { 850 return error.InvalidUpgradeHeader; 851 } 852 complete_response |= 2; 853 }, 854 10 => if (std.mem.eql(u8, line[0..i], "connection")) { 855 if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "upgrade")) { 856 return error.InvalidConnectionHeader; 857 } 858 complete_response |= 4; 859 }, 860 20 => if (std.mem.eql(u8, line[0..i], "sec-websocket-accept")) { 861 var h: [20]u8 = undefined; 862 { 863 var hasher = std.crypto.hash.Sha1.init(.{}); 864 hasher.update(key); 865 hasher.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); 866 hasher.final(&h); 867 } 868 869 var encoded_buf: [28]u8 = undefined; 870 const sec_hash = std.base64.standard.Encoder.encode(&encoded_buf, &h); 871 const header_value = std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace); 872 873 if (!std.mem.eql(u8, header_value, sec_hash)) { 874 return error.InvalidWebsocketAcceptHeader; 875 } 876 complete_response |= 8; 877 }, 878 24 => if (std.mem.eql(u8, line[0..i], "sec-websocket-extensions")) { 879 if (try parseExtension(line[i + 1 ..])) |sc| { 880 if (!compression) { 881 // server is saying compression, but we didn't ask for it. 882 return error.InvalidExtensionHeader; 883 } 884 if (!sc.client_no_context_takeover or !sc.server_no_context_takeover) { 885 // as of Zig 0.15, we no longer support context takeover 886 // We told the server this, it should have respected it. 887 return error.InvalidExtensionHeader; 888 } 889 890 server_compression = true; 891 } 892 }, 893 else => {}, // some other header we don't care about 894 } 895 } 896 } 897 898 if (milliTimestamp(stream.io) > deadline) { 899 return error.Timeout; 900 } 901 902 if (pos == buf.len) { 903 return error.ResponseTooLarge; 904 } 905 } 906 } 907 908 pub fn parseExtension(value: []const u8) !?ServerHandshake.Compression { 909 var deflate = false; 910 var client_max_bits: u8 = 15; 911 var client_no_context_takeover = false; 912 var server_no_context_takeover = false; 913 914 var it = std.mem.splitScalar(u8, value, ';'); 915 while (it.next()) |param_| { 916 const param = std.mem.trim(u8, param_, &ascii.whitespace); 917 if (std.mem.eql(u8, param, "permessage-deflate")) { 918 deflate = true; 919 continue; 920 } 921 if (std.mem.eql(u8, param, "client_no_context_takeover")) { 922 client_no_context_takeover = true; 923 continue; 924 } 925 if (std.mem.eql(u8, param, "server_no_context_takeover")) { 926 server_no_context_takeover = true; 927 continue; 928 } 929 const client_max_window_bits = "client_max_window_bits="; 930 if (std.mem.startsWith(u8, param, client_max_window_bits)) { 931 client_max_bits = std.fmt.parseInt(u8, param[client_max_window_bits.len..], 10) catch { 932 return error.InvalidCompressionServerMaxBits; 933 }; 934 } 935 } 936 if (deflate == false) { 937 return null; 938 } 939 940 if (client_max_bits != 15) { 941 // We don't offer client window, so if the server asks for one, that's an error 942 return error.InvalidExtensionHeader; 943 } 944 945 return .{ 946 .client_no_context_takeover = client_no_context_takeover, 947 .server_no_context_takeover = server_no_context_takeover, 948 }; 949 } 950}; 951 952const t = @import("../t.zig"); 953test "Client: handshake" { 954 { 955 // empty response 956 var pair = t.SocketPair.init(.{}); 957 defer pair.deinit(); 958 try pair.clientWriteAll("\r\n\r\n"); 959 960 var client = testClient(&pair); 961 defer client.deinit(); 962 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 963 } 964 965 { 966 // invalid websocket response 967 var pair = t.SocketPair.init(.{}); 968 defer pair.deinit(); 969 try pair.clientWriteAll("HTTP/1.1 200 OK\r\n\r\n"); 970 971 var client = testClient(&pair); 972 defer client.deinit(); 973 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 974 } 975 976 { 977 // missing upgrade header 978 var pair = t.SocketPair.init(.{}); 979 defer pair.deinit(); 980 try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\n\r\n"); 981 982 var client = testClient(&pair); 983 defer client.deinit(); 984 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 985 } 986 987 { 988 // wrong upgrade header 989 var pair = t.SocketPair.init(.{}); 990 defer pair.deinit(); 991 try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: nope\r\n\r\n"); 992 993 var client = testClient(&pair); 994 defer client.deinit(); 995 try t.expectError(error.InvalidUpgradeHeader, client.handshake("/", .{})); 996 } 997 998 { 999 // missing connection header 1000 var pair = t.SocketPair.init(.{}); 1001 defer pair.deinit(); 1002 try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\n\r\n"); 1003 1004 var client = testClient(&pair); 1005 defer client.deinit(); 1006 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 1007 } 1008 1009 { 1010 // wrong connection header 1011 var pair = t.SocketPair.init(.{}); 1012 defer pair.deinit(); 1013 try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: something\r\n\r\n"); 1014 1015 var client = testClient(&pair); 1016 defer client.deinit(); 1017 try t.expectError(error.InvalidConnectionHeader, client.handshake("/", .{})); 1018 } 1019 1020 { 1021 // missing Sec-Websocket-Accept header 1022 var pair = t.SocketPair.init(.{}); 1023 defer pair.deinit(); 1024 try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\nConnection: upgrade\r\n\r\n"); 1025 1026 var client = testClient(&pair); 1027 defer client.deinit(); 1028 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 1029 } 1030 1031 { 1032 // wrong Sec-Websocket-Accept header 1033 var pair = t.SocketPair.init(.{}); 1034 defer pair.deinit(); 1035 try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: hack\r\n\r\n"); 1036 1037 var client = testClient(&pair); 1038 defer client.deinit(); 1039 try t.expectError(error.InvalidWebsocketAcceptHeader, client.handshake("/", .{})); 1040 } 1041 1042 { 1043 // ok for successful 1044 var pair = t.SocketPair.init(.{}); 1045 defer pair.deinit(); 1046 try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: C/0nmHhBztSRGR1CwL6Tf4ZjwpY=\r\n\r\n"); 1047 1048 var client = testClient(&pair); 1049 defer client.deinit(); 1050 try client.handshake("/", .{}); 1051 try t.expectEqual(0, client._reader.pos); 1052 } 1053 1054 { 1055 // ok for successful, with overread 1056 var pair = t.SocketPair.init(.{}); 1057 defer pair.deinit(); 1058 try pair.clientWriteAll("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"); 1059 1060 var client = testClient(&pair); 1061 defer client.deinit(); 1062 try client.handshake("/", .{}); 1063 try t.expectEqual(50, client._reader.pos); 1064 } 1065} 1066 1067test "Client: handshake with terminating CRLF split across reads" { 1068 // regression: when TCP delivers the final \r of the blank-line CRLF as the 1069 // last byte of a read and the \n arrives in a later read, the end-of-headers 1070 // branch computed `over_read = pos - (line_start + 2)` while pos == line_start 1071 // + 1, underflowing usize → integer-overflow panic. parser must wait for \n. 1072 const io = std.Options.debug_io; 1073 1074 // generateKey() is deterministic under test ({1..16}); this accept matches it. 1075 const bin_key = [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; 1076 var encoded_key: [24]u8 = undefined; 1077 const key = std.base64.standard.Encoder.encode(&encoded_key, &bin_key); 1078 1079 const headers = "HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: C/0nmHhBztSRGR1CwL6Tf4ZjwpY=\r\n"; 1080 const trailing = "Some Random Data Which is Part Of the Next Message"; 1081 1082 const ChunkStream = struct { 1083 io: Io, 1084 chunks: []const []const u8, 1085 idx: usize = 0, 1086 fn read(self: *@This(), buf: []u8) !usize { 1087 if (self.idx >= self.chunks.len) return 0; 1088 const c = self.chunks[self.idx]; 1089 std.debug.assert(c.len <= buf.len); 1090 @memcpy(buf[0..c.len], c); 1091 self.idx += 1; 1092 return c.len; 1093 } 1094 fn readTimeout(self: *const @This(), ms: u32) !void { 1095 _ = self; 1096 _ = ms; 1097 } 1098 }; 1099 1100 // first read ends on the terminating \r; the \n (and over-read) arrive next. 1101 var chunks = [_][]const u8{ headers ++ "\r", "\n" ++ trailing }; 1102 var stream = ChunkStream{ .io = io, .chunks = &chunks }; 1103 1104 var buf: [4096]u8 = undefined; 1105 const opts = Client.HandshakeOpts{}; 1106 const res = try HandShakeReply.read(&buf, key, &opts, false, &stream); 1107 try t.expectEqual(trailing.len, res.over_read); 1108 try t.expectSlice(u8, trailing, buf[0..res.over_read]); 1109} 1110 1111test "Client: setting a timeout on a dead socket does not abort the process" { 1112 // regression: std.posix.setsockopt maps BADF/NOTSOCK/INVAL/FAULT to 1113 // `unreachable`. Those are reachable on a *connection* socket — the peer 1114 // resets, or another thread closes the fd, between connect and the timeout 1115 // call. `unreachable` cannot be caught, so an upstream dropping the 1116 // connection aborted the whole process from inside the handshake instead of 1117 // returning an error the caller could reconnect on. 1118 const io = std.Options.debug_io; 1119 const stream = try testConnectedStream(io, "127.0.0.1", 9292); 1120 var dead = Stream.init(io, stream, null); 1121 // Close underneath the Stream: every later setsockopt sees EBADF, which is 1122 // precisely the race the stdlib declares impossible. 1123 stream.close(io); 1124 1125 // Before the fix each of these panicked rather than returning. 1126 try dead.readTimeout(5_000); 1127 try dead.writeTimeout(5_000); 1128 try dead.readTimeout(0); 1129} 1130 1131test "Client: write/read" { 1132 const io = std.Options.debug_io; 1133 const stream = try testConnectedStream(io, "127.0.0.1", 9292); 1134 var client = try Client.initWithStream(io, t.allocator, stream, .{ 1135 .port = 9292, 1136 .host = "127.0.0.1", 1137 }); 1138 defer client.deinit(); 1139 1140 try client.handshake("/", .{ 1141 .timeout_ms = 1000, 1142 }); 1143 1144 var buf = [_]u8{ 'o', 'v', 'e', 'r' }; 1145 try client.write(&buf); 1146 try client.readTimeout(1000); 1147 1148 const message = (try client.read()) orelse unreachable; 1149 try t.expectEqual(.text, message.type); 1150 try t.expectString("9000", message.data); 1151 1152 client.close(.{}) catch unreachable; 1153} 1154 1155test "Client: close with code" { 1156 const io = std.Options.debug_io; 1157 const stream = try testConnectedStream(io, "127.0.0.1", 9292); 1158 var client = try Client.initWithStream(io, t.allocator, stream, .{ 1159 .port = 9292, 1160 .host = "127.0.0.1", 1161 }); 1162 defer client.deinit(); 1163 1164 try client.handshake("/", .{ 1165 .timeout_ms = 1000, 1166 }); 1167 1168 client.close(.{ .code = 4002 }) catch unreachable; 1169} 1170 1171test "Client: with code and reason" { 1172 const io = std.Options.debug_io; 1173 const stream = try testConnectedStream(io, "127.0.0.1", 9292); 1174 var client = try Client.initWithStream(io, t.allocator, stream, .{ 1175 .port = 9292, 1176 .host = "127.0.0.1", 1177 }); 1178 defer client.deinit(); 1179 1180 try client.handshake("/", .{ 1181 .timeout_ms = 1000, 1182 }); 1183 1184 client.close(.{ .code = 4002, .reason = "goodbye" }) catch unreachable; 1185} 1186 1187test "Client: Handler" { 1188 var h = try ClientHandler.init(t.allocator); 1189 defer h.deinit(); 1190 1191 var buf: [6]u8 = undefined; 1192 { 1193 @memcpy(buf[0..3], "dyn"); 1194 try h.client.write(buf[0..3]); 1195 } 1196 1197 { 1198 @memcpy(buf[0..4], "ping"); 1199 try h.client.write(buf[0..4]); 1200 } 1201 1202 { 1203 @memcpy(buf[0..4], "pong"); 1204 try h.client.write(buf[0..4]); 1205 } 1206 1207 { 1208 @memcpy(buf[0..6], "close1"); 1209 try h.client.write(buf[0..6]); 1210 } 1211 1212 try h.client.readLoop(&h); 1213 1214 // if pong is true then ping and message have to be true 1215 // because each asserts the previous 1216 try t.expectEqual(true, h.pong); 1217 try t.expectEqual(true, h.closed); 1218} 1219 1220fn testClient(pair: *t.SocketPair) Client { 1221 pair.server_taken = true; 1222 const stream = pair.server; 1223 const io = std.Options.debug_io; 1224 const bp = t.allocator.create(buffer.Provider) catch unreachable; 1225 bp.* = buffer.Provider.init(t.allocator, .{ .count = 0, .size = 0, .max = 4096 }) catch unreachable; 1226 1227 const reader_buf = bp.allocator.alloc(u8, 1024) catch unreachable; 1228 1229 return .{ 1230 .io = io, 1231 ._closed = false, 1232 ._own_bp = true, 1233 ._mask_fn = generateMask, 1234 ._compression_opts = null, 1235 .stream = .{ .io = io, .stream = stream }, 1236 ._reader = Reader.init(reader_buf, bp, null), 1237 }; 1238} 1239 1240fn testConnectedStream(io: Io, host: []const u8, port: u16) !net.Stream { 1241 const addr = try net.IpAddress.parse(host, port); 1242 return net.IpAddress.connect(&addr, io, .{ .mode = .stream }); 1243} 1244 1245const ClientHandler = struct { 1246 ping: bool = false, 1247 pong: bool = false, 1248 closed: bool = false, 1249 message: bool = false, 1250 client: Client, 1251 1252 fn init(allocator: Allocator) !ClientHandler { 1253 const io = std.Options.debug_io; 1254 const stream = try testConnectedStream(io, "127.0.0.1", 9292); 1255 var client = try Client.initWithStream(io, allocator, stream, .{ 1256 .port = 9292, 1257 .host = "127.0.0.1", 1258 }); 1259 errdefer client.deinit(); 1260 1261 try client.handshake("/", .{ 1262 .timeout_ms = 1000, 1263 }); 1264 1265 return .{ 1266 .client = client, 1267 }; 1268 } 1269 1270 fn deinit(self: *ClientHandler) void { 1271 self.client.deinit(); 1272 } 1273 1274 pub fn serverMessage(self: *ClientHandler, data: []u8, tpe: proto.Message.TextType) !void { 1275 try t.expectEqual(.text, tpe); 1276 try t.expectString("over 9000!", data); 1277 self.message = true; 1278 } 1279 1280 pub fn serverPing(self: *ClientHandler, data: []u8) !void { 1281 try t.expectEqual(true, self.message); 1282 try t.expectString("a-ping", data); 1283 self.ping = true; 1284 } 1285 1286 pub fn serverPong(self: *ClientHandler, data: []u8) !void { 1287 try t.expectEqual(true, self.ping); 1288 try t.expectString("a-pong", data); 1289 self.pong = true; 1290 } 1291 1292 pub fn close(self: *ClientHandler) void { 1293 self.client.close(.{}) catch unreachable; 1294 self.closed = true; 1295 } 1296};