A websocket implementation for zig
0

Configure Feed

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

websocket.zig / src / client / client.zig
40 kB 1119 lines
1const std = @import("std"); 2 3const posix = @import("../posix.zig"); 4const proto = @import("../proto.zig"); 5const buffer = @import("../buffer.zig"); 6const CompressionOpts = @import("../websocket.zig").Compression; 7const ServerHandshake = @import("../server/handshake.zig").Handshake; 8 9const Io = std.Io; 10const ascii = std.ascii; 11const tls = std.crypto.tls; 12const log = std.log.scoped(.websocket); 13 14const Reader = proto.Reader; 15const Allocator = std.mem.Allocator; 16const Bundle = std.crypto.Certificate.Bundle; 17 18fn ReadLoopHandler(comptime T: type) type { 19 const info = @typeInfo(T); 20 21 switch (info) { 22 .@"struct" => |struct_info| { 23 if (struct_info.is_tuple) 24 @compileError("readLoop: handler does not support tuples."); 25 26 return T; 27 }, 28 .pointer => |ptr_info| { 29 switch (ptr_info.size) { 30 .one => return ReadLoopHandler(ptr_info.child), 31 else => @compileError("readLoop: handler does not support Slice, C and Many pointers."), 32 } 33 }, 34 else => @compileError("readLoop: expected handler to be a struct or pointer to a struct but found '" ++ @tagName(info) ++ "'"), 35 } 36} 37 38pub const Client = struct { 39 io: Io, 40 stream: Stream, 41 _reader: Reader, 42 _closed: bool, 43 _compression_opts: ?CompressionOpts, 44 _compression: ?Client.Compression = null, 45 46 // When creating a client, we can either be given a BufferProvider or create 47 // one ourselves. If we create it ourselves (in init), we "own" it and must 48 // free it on deinit. (The reference to the buffer provider is already in the 49 // reader, no need to hold another reference in the client). 50 _own_bp: bool, 51 52 // For advanced cases, a custom masking function can be provided. Masking 53 // is a security feature that only really makes sense in the browser. If you 54 // aren't running websockets in the browser AND you control both the client 55 // and the server, you could get a performance boost by not masking. 56 _mask_fn: *const fn (Io) [4]u8, 57 58 pub const Config = struct { 59 port: u16, 60 host: []const u8, 61 tls: bool = false, 62 max_size: usize = 65536, 63 buffer_size: usize = 4096, 64 ca_bundle: ?Bundle = null, 65 mask_fn: *const fn (Io) [4]u8 = generateMask, 66 buffer_provider: ?*buffer.Provider = null, 67 compression: ?CompressionOpts = null, 68 }; 69 70 pub const HandshakeOpts = struct { 71 timeout_ms: u32 = 10000, 72 headers: ?[]const u8 = null, 73 }; 74 75 const Compression = struct { 76 allocator: Allocator, 77 retain_writer: bool, 78 write_treshold: usize, 79 writer: Io.Writer.Allocating, 80 }; 81 82 pub fn init(io: Io, allocator: Allocator, config: Config) !Client { 83 if (config.compression != null) { 84 log.err("Compression is disabled as part of the 0.15 upgrade. I do hope to re-enable it soon.", .{}); 85 return error.InvalidConfiguraion; 86 } 87 88 const host_name = try Io.net.HostName.init(config.host); 89 const net_stream = try host_name.connect(io, config.port, .{ .mode = .stream }); 90 91 var tls_client: ?*TLSClient = null; 92 if (config.tls) { 93 tls_client = try TLSClient.init(io, allocator, net_stream, &config); 94 } 95 const stream = Stream.init(io, net_stream, tls_client); 96 97 var own_bp = false; 98 var buffer_provider: *buffer.Provider = undefined; 99 100 // If a buffer_provider is provided, we'll use that. 101 // If it isn't, we need to create one which also means we now "own" it 102 // and we're responsible for cleaning it up 103 if (config.buffer_provider) |shared_bp| { 104 buffer_provider = shared_bp; 105 } else { 106 own_bp = true; 107 buffer_provider = try allocator.create(buffer.Provider); 108 errdefer allocator.destroy(buffer_provider); 109 buffer_provider.* = try buffer.Provider.init(io, allocator, .{ 110 .size = 0, 111 .count = 0, 112 .max = config.max_size, 113 }); 114 } 115 116 errdefer if (own_bp) { 117 buffer_provider.deinit(); 118 allocator.destroy(buffer_provider); 119 }; 120 121 const reader_buf = try buffer_provider.allocator.alloc(u8, config.buffer_size); 122 errdefer buffer_provider.allocator.free(reader_buf); 123 124 return .{ 125 .io = io, 126 .stream = stream, 127 ._closed = false, 128 ._own_bp = own_bp, 129 ._mask_fn = config.mask_fn, 130 ._compression_opts = null, //TODO: ZIG 0.15 131 ._reader = Reader.init(reader_buf, buffer_provider, null), 132 }; 133 } 134 135 pub fn deinit(self: *Client) void { 136 self.closeStream(); 137 138 const larger_buffer_provider = self._reader.large_buffer_provider; 139 const allocator = larger_buffer_provider.allocator; 140 allocator.free(self._reader.static); 141 142 self._reader.deinit(); 143 144 if (self._own_bp) { 145 larger_buffer_provider.deinit(); 146 allocator.destroy(larger_buffer_provider); 147 } 148 } 149 150 pub fn handshake(self: *Client, path: []const u8, opts: HandshakeOpts) !void { 151 const stream = &self.stream; 152 errdefer self.closeStream(); 153 154 // we've already setup our reader, and the reader has a static buffer 155 // we might as well use it! 156 const buf = self._reader.static; 157 const key = blk: { 158 const bin_key = generateKey(self.io); 159 var encoded_key: [24]u8 = undefined; 160 break :blk std.base64.standard.Encoder.encode(&encoded_key, &bin_key); 161 }; 162 163 try sendHandshake(path, key, buf, &opts, self._compression_opts != null, stream); 164 165 const res = try HandShakeReply.read(self.io, buf, key, &opts, self._compression_opts != null, stream); 166 errdefer self.close(.{ .code = 1001 }) catch unreachable; 167 168 // Set up compression with agreed-on parameters 169 if (res.compression) { 170 try self.setupCompression(); 171 } 172 173 // We might have read more than handshake response. If so, readHandshakeReply 174 // has positioned the extra data at the start of the buffer, but we need 175 // to set the length. 176 self._reader.pos = res.over_read; 177 } 178 179 fn setupCompression(self: *Client) !void { 180 std.debug.assert(self._compression_opts != null); 181 self._reader.allow_compressed = true; 182 183 const allocator = self._reader.large_buffer_provider.allocator; 184 const config = self._compression_opts.?; 185 self._compression = .{ 186 .allocator = allocator, 187 .write_treshold = config.write_threshold.?, 188 .retain_writer = config.retain_write_buffer, 189 .writer = std.Io.Writer.Allocating.init(allocator), 190 }; 191 } 192 193 pub fn readLoop(self: *Client, handler: anytype) !void { 194 const Handler = ReadLoopHandler(@TypeOf(handler)); 195 var reader = &self._reader; 196 197 defer if (comptime std.meta.hasFn(Handler, "close")) { 198 handler.close(); 199 }; 200 201 // block until we have data 202 try self.readTimeout(0); 203 204 while (true) { 205 const message = self.read() catch |err| switch (err) { 206 error.Closed => return, 207 else => return err, 208 } orelse unreachable; 209 210 const message_type = message.type; 211 defer reader.done(message_type); 212 213 switch (message_type) { 214 .text, .binary => { 215 switch (comptime @typeInfo(@TypeOf(Handler.serverMessage)).@"fn".params.len) { 216 2 => try handler.serverMessage(message.data), 217 3 => try handler.serverMessage(message.data, if (message_type == .text) .text else .binary), 218 else => @compileError(@typeName(Handler) ++ ".serverMessage must accept 2 or 3 parameters"), 219 } 220 }, 221 .ping => if (comptime std.meta.hasFn(Handler, "serverPing")) { 222 try handler.serverPing(message.data); 223 } else { 224 // @constCast is safe because we know message.data points to 225 // reader.buffer.buf, which we own and which can be mutated 226 try self.writeFrame(.pong, @constCast(message.data)); 227 }, 228 .close => { 229 if (comptime std.meta.hasFn(Handler, "serverClose")) { 230 try handler.serverClose(message.data); 231 } else { 232 self.close(.{}) catch unreachable; 233 } 234 return; 235 }, 236 .pong => if (comptime std.meta.hasFn(Handler, "serverPong")) { 237 try handler.serverPong(message.data); 238 }, 239 } 240 } 241 } 242 243 pub fn read(self: *Client) !?proto.Message { 244 var reader = &self._reader; 245 const stream = &self.stream; 246 247 while (true) { 248 // try to read a message from our buffer first, before trying to 249 // get more data from the socket. 250 const has_more, const message = reader.read() catch |err| { 251 self.close(.{ .code = 1002 }) catch unreachable; 252 return err; 253 } orelse { 254 reader.fill(stream) catch |err| switch (err) { 255 error.WouldBlock => return null, 256 error.Closed, error.ConnectionResetByPeer, error.BrokenPipe, error.NotOpenForReading => { 257 @atomicStore(bool, &self._closed, true, .monotonic); 258 return error.Closed; 259 }, 260 else => { 261 self.close(.{ .code = 1002 }) catch unreachable; 262 return err; 263 }, 264 }; 265 continue; 266 }; 267 268 _ = has_more; 269 return message; 270 } 271 } 272 273 pub fn done(self: *Client, message: proto.Message) void { 274 self._reader.done(message.type); 275 } 276 277 pub fn readLoopInNewThread(self: *Client, h: anytype) !std.Thread { 278 return std.Thread.spawn(.{}, readLoopOwnedThread, .{ self, h }); 279 } 280 281 fn readLoopOwnedThread(self: *Client, h: anytype) void { 282 self.readLoop(h) catch {}; 283 } 284 285 pub fn writeTimeout(self: *const Client, ms: u32) !void { 286 return self.stream.writeTimeout(ms); 287 } 288 289 pub fn readTimeout(self: *Client, ms: u32) !void { 290 return self.stream.readTimeout(ms); 291 } 292 293 pub fn write(self: *Client, data: []u8) !void { 294 return self.writeFrame(.text, data); 295 } 296 297 pub fn writeText(self: *Client, data: []u8) !void { 298 return self.writeFrame(.text, data); 299 } 300 301 pub fn writeBin(self: *Client, data: []u8) !void { 302 return self.writeFrame(.binary, data); 303 } 304 305 pub fn writePing(self: *Client, data: []u8) !void { 306 return self.writeFrame(.ping, data); 307 } 308 309 pub fn writePong(self: *Client, data: []u8) !void { 310 return self.writeFrame(.pong, data); 311 } 312 313 const CloseOpts = struct { 314 code: ?u16 = null, 315 reason: []const u8 = "", 316 }; 317 318 pub fn close(self: *Client, opts: CloseOpts) !void { 319 if (@atomicRmw(bool, &self._closed, .Xchg, true, .monotonic) == true) { 320 // already closed 321 return; 322 } 323 324 defer self.stream.close(); 325 326 const code = opts.code orelse { 327 self.writeFrame(.close, "") catch {}; 328 return; 329 }; 330 331 const reason = opts.reason; 332 if (reason.len > 123) { 333 return error.ReasonTooLong; 334 } 335 336 var buf: [125]u8 = undefined; 337 buf[0] = @intCast((code >> 8) & 0xFF); 338 buf[1] = @intCast(code & 0xFF); 339 340 const end = 2 + reason.len; 341 @memcpy(buf[2..end], reason); 342 self.writeFrame(.close, buf[0..end]) catch {}; 343 } 344 345 pub fn writeFrame(self: *Client, op_code: proto.OpCode, data: []u8) !void { 346 const payload = data; 347 const compressed = false; 348 // if (self._compression) |c| { 349 // if (data.len >= c.write_treshold and (op_code == .binary or op_code == .text)) { 350 // compressed = true; 351 352 // var writer = &c.writer; 353 // var compressor = &c.compressor; 354 // var fbs = std.io.fixedBufferStream(data); 355 // _ = try compressor.compress(fbs.reader()); 356 // try compressor.flush(); 357 // payload = writer.items[0 .. writer.items.len - 4]; 358 359 // if (c.reset) { 360 // c.compressor = try Compression.Type.init(writer.writer(), .{}); 361 // } 362 // } 363 // } 364 // defer if (compressed) { 365 // const c = self._compression.?; 366 // if (c.retain_writer) { 367 // c.compressor.wrt.context.clearRetainingCapacity(); 368 // } else { 369 // c.compressor.wrt.context.clearAndFree(); 370 // } 371 // }; 372 373 // maximum possible prefix length. op_code + length_type + 8byte length + 4 byte mask 374 var buf: [14]u8 = undefined; 375 const header = proto.writeFrameHeader(&buf, op_code, payload.len, compressed); 376 377 const header_len = header.len; 378 const header_end = header.len + 4; // for the mask 379 380 buf[1] |= 128; // indicate that the payload is masked 381 382 const mask = self._mask_fn(self.io); 383 @memcpy(buf[header_len..header_end], &mask); 384 try self.stream.writeAll(buf[0..header_end]); 385 386 if (payload.len > 0) { 387 proto.mask(&mask, payload); 388 try self.stream.writeAll(payload); 389 } 390 } 391 392 fn closeStream(self: *Client) void { 393 if (@atomicRmw(bool, &self._closed, .Xchg, true, .monotonic) == false) { 394 self.stream.close(); 395 } 396 } 397}; 398 399pub const Stream = struct { 400 io: Io, 401 stream: Io.net.Stream, 402 tls_client: ?*TLSClient = null, 403 read_timeout_ms: u32 = 0, 404 405 pub fn init(io: Io, stream: Io.net.Stream, tls_client: ?*TLSClient) Stream { 406 return .{ 407 .io = io, 408 .stream = stream, 409 .tls_client = tls_client, 410 }; 411 } 412 413 pub fn close(self: *Stream) void { 414 const fd = self.stream.socket.handle; 415 const builtin = @import("builtin"); 416 const native_os = builtin.os.tag; 417 418 if (self.tls_client) |tls_client| { 419 // Shutdown the socket first, so readLoop() can exit, before tls_client's buffers are freed 420 if (native_os == .wasi and !builtin.link_libc) { 421 _ = std.os.wasi.sock_shutdown(fd, .{ .WR = true, .RD = true }); 422 } else { 423 posix.shutdown(fd, .both) catch {}; 424 } 425 tls_client.deinit(); 426 } 427 428 // posix.close panics on EBADF 429 // This is a general issue in Zig: 430 // https://github.com/ziglang/zig/issues/6389 431 // 432 // we don't want to crash on double close 433 434 if (native_os == .windows) { 435 return std.os.windows.CloseHandle(fd); 436 } 437 if (native_os == .wasi and !builtin.link_libc) { 438 _ = std.os.wasi.fd_close(fd); 439 return; 440 } 441 _ = std.posix.system.close(fd); 442 } 443 444 pub fn read(self: *Stream, buf: []u8) !usize { 445 // A read timeout is implemented by polling the socket for readiness 446 // rather than SO_RCVTIMEO. On the TLS path the underlying read goes 447 // through std.crypto.tls, whose reader treats a socket EAGAIN (what 448 // SO_RCVTIMEO produces on timeout) as a programmer bug and panics in 449 // debug builds. Polling lets a timeout surface as error.WouldBlock (which 450 // read() turns into "no message") without ever issuing a read that could 451 // EAGAIN. 452 // 453 // Skip the poll when the TLS client already has decrypted plaintext 454 // buffered: a previous read can decrypt more than the caller consumed, so 455 // the socket can be empty (poll would time out) even though data is 456 // available in-process, which would starve it. 457 const tls_buffered = if (self.tls_client) |tls_client| tls_client.client.reader.bufferedLen() else 0; 458 if (self.read_timeout_ms > 0 and tls_buffered == 0) { 459 var pfd = [_]std.posix.pollfd{.{ 460 .fd = self.stream.socket.handle, 461 .events = std.posix.POLL.IN, 462 .revents = 0, 463 }}; 464 // A poll failure is a real read failure, not "no data": surface it 465 // (mapped into this read path's error set) rather than swallowing it. 466 const ready = std.posix.poll(&pfd, @intCast(self.read_timeout_ms)) catch return error.ReadFailed; 467 if (ready == 0) return error.WouldBlock; 468 } 469 if (self.tls_client) |tls_client| { 470 var w: std.Io.Writer = .fixed(buf); 471 while (true) { 472 const n = try tls_client.client.reader.stream(&w, .limited(buf.len)); 473 if (n != 0) { 474 return n; 475 } 476 } 477 } 478 return posix.read(self.stream.socket.handle, buf); 479 } 480 481 pub fn writeAll(self: *Stream, data: []const u8) !void { 482 if (self.tls_client) |tls_client| { 483 try tls_client.client.writer.writeAll(data); 484 // I know this looks silly, but as far as I can tell, this is what 485 // we need to do. 486 try tls_client.client.writer.flush(); 487 try tls_client.stream_writer.interface.flush(); 488 return; 489 } 490 491 var writer = self.stream.writer(self.io, &.{}); 492 try writer.interface.writeAll(data); 493 return writer.interface.flush(); 494 } 495 496 const zero_timeout = std.mem.toBytes(posix.timeval{ .sec = 0, .usec = 0 }); 497 pub fn writeTimeout(self: *const Stream, ms: u32) !void { 498 return self.setTimeout(posix.SO.SNDTIMEO, ms); 499 } 500 501 pub fn readTimeout(self: *Stream, ms: u32) !void { 502 // Stored and applied via poll() in read(); see the note there for why this 503 // does not use SO_RCVTIMEO. 504 self.read_timeout_ms = ms; 505 } 506 507 fn setTimeout(self: *const Stream, opt_name: u32, ms: u32) !void { 508 if (ms == 0) { 509 return self.setsockopt(opt_name, &zero_timeout); 510 } 511 512 const timeout = std.mem.toBytes(posix.timeval{ 513 .sec = @intCast(@divTrunc(ms, 1000)), 514 .usec = @intCast(@mod(ms, 1000) * 1000), 515 }); 516 return self.setsockopt(opt_name, &timeout); 517 } 518 519 pub fn setsockopt(self: *const Stream, opt_name: u32, value: []const u8) !void { 520 return posix.setsockopt(self.stream.socket.handle, posix.SOL.SOCKET, opt_name, value); 521 } 522}; 523 524const TLSClient = struct { 525 io: Io, 526 client: tls.Client, 527 stream: Io.net.Stream, 528 stream_writer: Io.net.Stream.Writer, 529 stream_reader: Io.net.Stream.Reader, 530 arena: std.heap.ArenaAllocator, 531 532 fn init(io: Io, allocator: Allocator, stream: Io.net.Stream, config: *const Client.Config) !*TLSClient { 533 var arena = std.heap.ArenaAllocator.init(allocator); 534 errdefer arena.deinit(); 535 536 const aa = arena.allocator(); 537 538 // 0.16: Bundle is heap-allocated so we can pass a pointer to TLS 539 // Options.ca.bundle. A single-threaded RwLock is fine here because 540 // the bundle is only touched by this TLS client; the RwLock serves 541 // only to match the Options.ca.bundle contract. 542 const bundle_ptr = try aa.create(Bundle); 543 if (config.ca_bundle) |existing| { 544 bundle_ptr.* = existing; 545 } else { 546 bundle_ptr.* = .empty; 547 // 0.16: rescan signature is (*Bundle, gpa, io, now: Io.Timestamp). 548 try bundle_ptr.rescan(aa, io, Io.Timestamp.now(io, .real)); 549 } 550 const bundle_lock = try aa.create(Io.RwLock); 551 bundle_lock.* = .init; 552 553 // The TLS input and output have to be max_ciphertext_record_len each. 554 // It isn't clear to me how big the un-encrypted reader and writer 555 // need to be. I would think 0, but that will fail an assertion. I 556 // don't think that it's right that we need 4 buffers, but apparently 557 // we do. Until i figure this out, using 4 x max_ciphertext_record_len 558 // seems like the only safe choice. 559 const buf_len = std.crypto.tls.max_ciphertext_record_len; 560 var buf = try aa.alloc(u8, buf_len * 4); 561 562 const self = try aa.create(TLSClient); 563 self.* = .{ 564 .io = io, 565 .stream = stream, 566 .arena = arena, 567 .client = undefined, 568 .stream_writer = stream.writer(io, buf.ptr[0..buf_len][0..buf_len]), 569 .stream_reader = stream.reader(io, buf.ptr[buf_len .. 2 * buf_len][0..buf_len]), 570 }; 571 572 // 0.16 TLS Client.Options requires `entropy` and `realtime_now` in 573 // addition to the 0.15 set. Fill both from the shim Io — the 574 // entropy buffer is read only during `init`. 575 var entropy_buf: [tls.Client.Options.entropy_len]u8 = undefined; 576 io.random(&entropy_buf); 577 578 self.client = try tls.Client.init( 579 &self.stream_reader.interface, 580 &self.stream_writer.interface, 581 .{ 582 .ca = .{ .bundle = .{ 583 .gpa = aa, 584 .io = io, 585 .lock = bundle_lock, 586 .bundle = bundle_ptr, 587 } }, 588 .host = .{ .explicit = config.host }, 589 .read_buffer = buf.ptr[2 * buf_len .. 3 * buf_len][0..buf_len], 590 .write_buffer = buf.ptr[3 * buf_len .. 4 * buf_len][0..buf_len], 591 .entropy = &entropy_buf, 592 .realtime_now = std.Io.Timestamp.now(io, .real), 593 }, 594 ); 595 596 return self; 597 } 598 599 fn deinit(self: *TLSClient) void { 600 _ = self.client.end() catch {}; 601 self.arena.deinit(); 602 } 603}; 604 605fn generateKey(io: Io) [16]u8 { 606 if (comptime @import("builtin").is_test) { 607 return [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; 608 } 609 var key: [16]u8 = undefined; 610 io.random(&key); 611 return key; 612} 613 614fn generateMask(io: Io) [4]u8 { 615 var m: [4]u8 = undefined; 616 io.random(&m); 617 return m; 618} 619 620fn sendHandshake(path: []const u8, key: []const u8, buf: []u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !void { 621 @memcpy(buf[0..4], "GET "); 622 var pos: usize = 4; 623 var end = pos + path.len; 624 625 { 626 @memcpy(buf[pos..end], path); 627 pos = end; 628 } 629 630 { 631 const headers = " HTTP/1.1\r\ncontent-length: 0\r\nupgrade: websocket\r\nsec-websocket-version: 13\r\nconnection: upgrade\r\nsec-websocket-key: "; 632 end = pos + headers.len; 633 @memcpy(buf[pos..end], headers); 634 635 pos = end; 636 end = pos + key.len; 637 @memcpy(buf[pos..end], key); 638 } 639 640 if (compression) { 641 // NOTE: client_max_window_bits is unsupported 642 const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover"; 643 pos = end; 644 end = pos + permessage_deflate.len; 645 @memcpy(buf[pos..end], permessage_deflate); 646 } 647 648 { 649 pos = end; 650 end = pos + 2; 651 @memcpy(buf[pos..end], "\r\n"); 652 pos = end; 653 } 654 655 if (opts.headers) |extra_headers| { 656 end = pos + extra_headers.len; 657 @memcpy(buf[pos..end], extra_headers); 658 pos = end; 659 if (!std.mem.endsWith(u8, extra_headers, "\r\n")) { 660 buf[pos] = '\r'; 661 buf[pos + 1] = '\n'; 662 pos += 2; 663 } 664 } 665 buf[pos] = '\r'; 666 buf[pos + 1] = '\n'; 667 668 try stream.writeTimeout(opts.timeout_ms); 669 try stream.writeAll(buf[0 .. pos + 2]); 670 try stream.writeTimeout(0); 671} 672 673const HandShakeReply = struct { 674 compression: bool, 675 over_read: usize, 676 677 fn read(io: Io, buf: []u8, key: []const u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !HandShakeReply { 678 const timeout_ms = opts.timeout_ms; 679 // 0.16 removed `std.time.milliTimestamp`; compute ms since epoch 680 // from `std.Io.Timestamp.now(io, .real)` (nanoseconds). 681 const deadline = @divTrunc(Io.Timestamp.now(io, .real).nanoseconds, std.time.ns_per_ms) + timeout_ms; 682 try stream.readTimeout(timeout_ms); 683 684 var pos: usize = 0; 685 var line_start: usize = 0; 686 var complete_response: u8 = 0; 687 var server_compression: bool = false; 688 689 while (true) { 690 const n = stream.read(buf[pos..]) catch |err| { 691 // `error.WouldBlock` may not be in `err`'s set on Windows 692 // (where the read goes through ReadFile), so match by name. 693 if (std.mem.eql(u8, @errorName(err), "WouldBlock")) return error.Timeout; 694 return err; 695 }; 696 if (n == 0) { 697 return error.ConnectionClosed; 698 } 699 700 pos += n; 701 while (std.mem.indexOfScalar(u8, buf[line_start..pos], '\r')) |relative_end| { 702 if (relative_end == 0) { 703 if (complete_response != 15) { 704 return error.InvalidHandshakeResponse; 705 } 706 const over_read = pos - (line_start + 2); 707 std.mem.copyForwards(u8, buf[0..over_read], buf[line_start + 2 .. pos]); 708 try stream.readTimeout(0); 709 return .{ 710 .over_read = over_read, 711 .compression = server_compression, 712 }; 713 } 714 715 const line_end = line_start + relative_end; 716 const line = buf[line_start..line_end]; 717 718 // the next line starts where this line ends, skip over the \r\n 719 line_start = line_end + 2; 720 721 if (complete_response == 0) { 722 if (!ascii.startsWithIgnoreCase(line, "HTTP/1.1 101 ")) { 723 return error.InvalidHandshakeResponse; 724 } 725 complete_response |= 1; 726 continue; 727 } 728 729 for (line, 0..) |b, i| { 730 // find the colon and lowercase the header while we're iterating 731 if ('A' <= b and b <= 'Z') { 732 line[i] = b + 32; 733 continue; 734 } 735 736 if (b != ':') { 737 continue; 738 } 739 740 switch (i) { 741 7 => if (std.mem.eql(u8, line[0..i], "upgrade")) { 742 if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "websocket")) { 743 return error.InvalidUpgradeHeader; 744 } 745 complete_response |= 2; 746 }, 747 10 => if (std.mem.eql(u8, line[0..i], "connection")) { 748 if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "upgrade")) { 749 return error.InvalidConnectionHeader; 750 } 751 complete_response |= 4; 752 }, 753 20 => if (std.mem.eql(u8, line[0..i], "sec-websocket-accept")) { 754 var h: [20]u8 = undefined; 755 { 756 var hasher = std.crypto.hash.Sha1.init(.{}); 757 hasher.update(key); 758 hasher.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); 759 hasher.final(&h); 760 } 761 762 var encoded_buf: [28]u8 = undefined; 763 const sec_hash = std.base64.standard.Encoder.encode(&encoded_buf, &h); 764 const header_value = std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace); 765 766 if (!std.mem.eql(u8, header_value, sec_hash)) { 767 return error.InvalidWebsocketAcceptHeader; 768 } 769 complete_response |= 8; 770 }, 771 24 => if (std.mem.eql(u8, line[0..i], "sec-websocket-extensions")) { 772 if (try parseExtension(line[i + 1 ..])) |sc| { 773 if (!compression) { 774 // server is saying compression, but we didn't ask for it. 775 return error.InvalidExtensionHeader; 776 } 777 if (!sc.client_no_context_takeover or !sc.server_no_context_takeover) { 778 // as of Zig 0.15, we no longer support context takeover 779 // We told the server this, it should have respected it. 780 return error.InvalidExtensionHeader; 781 } 782 783 server_compression = true; 784 } 785 }, 786 else => {}, // some other header we don't care about 787 } 788 } 789 } 790 791 if (@divTrunc(Io.Timestamp.now(io, .real).nanoseconds, std.time.ns_per_ms) > deadline) { 792 return error.Timeout; 793 } 794 795 if (pos == buf.len) { 796 return error.ResponseTooLarge; 797 } 798 } 799 } 800 801 pub fn parseExtension(value: []const u8) !?ServerHandshake.Compression { 802 var deflate = false; 803 var client_max_bits: u8 = 15; 804 var client_no_context_takeover = false; 805 var server_no_context_takeover = false; 806 807 var it = std.mem.splitScalar(u8, value, ';'); 808 while (it.next()) |param_| { 809 const param = std.mem.trim(u8, param_, &ascii.whitespace); 810 if (std.mem.eql(u8, param, "permessage-deflate")) { 811 deflate = true; 812 continue; 813 } 814 if (std.mem.eql(u8, param, "client_no_context_takeover")) { 815 client_no_context_takeover = true; 816 continue; 817 } 818 if (std.mem.eql(u8, param, "server_no_context_takeover")) { 819 server_no_context_takeover = true; 820 continue; 821 } 822 const client_max_window_bits = "client_max_window_bits="; 823 if (std.mem.startsWith(u8, param, client_max_window_bits)) { 824 client_max_bits = std.fmt.parseInt(u8, param[client_max_window_bits.len..], 10) catch { 825 return error.InvalidCompressionServerMaxBits; 826 }; 827 } 828 } 829 if (deflate == false) { 830 return null; 831 } 832 833 if (client_max_bits != 15) { 834 // We don't offer client window, so if the server asks for one, that's an error 835 return error.InvalidExtensionHeader; 836 } 837 838 return .{ 839 .client_no_context_takeover = client_no_context_takeover, 840 .server_no_context_takeover = server_no_context_takeover, 841 }; 842 } 843}; 844 845const t = @import("../t.zig"); 846test "Client: handshake" { 847 { 848 // empty response 849 var pair = t.SocketPair.init(.{}); 850 defer pair.deinit(); 851 var writer = pair.client.writer(t.io, &.{}); 852 try writer.interface.writeAll("\r\n\r\n"); 853 854 var client = testClient(pair.server); 855 defer client.deinit(); 856 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 857 } 858 859 { 860 // invalid websocket response 861 var pair = t.SocketPair.init(.{}); 862 defer pair.deinit(); 863 var writer = pair.client.writer(t.io, &.{}); 864 try writer.interface.writeAll("HTTP/1.1 200 OK\r\n\r\n"); 865 866 var client = testClient(pair.server); 867 defer client.deinit(); 868 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 869 } 870 871 { 872 // missing upgrade header 873 var pair = t.SocketPair.init(.{}); 874 defer pair.deinit(); 875 var writer = pair.client.writer(t.io, &.{}); 876 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\n\r\n"); 877 878 var client = testClient(pair.server); 879 defer client.deinit(); 880 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 881 } 882 883 { 884 // wrong upgrade header 885 var pair = t.SocketPair.init(.{}); 886 defer pair.deinit(); 887 var writer = pair.client.writer(t.io, &.{}); 888 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: nope\r\n\r\n"); 889 890 var client = testClient(pair.server); 891 defer client.deinit(); 892 try t.expectError(error.InvalidUpgradeHeader, client.handshake("/", .{})); 893 } 894 895 { 896 // missing connection header 897 var pair = t.SocketPair.init(.{}); 898 defer pair.deinit(); 899 var writer = pair.client.writer(t.io, &.{}); 900 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\n\r\n"); 901 902 var client = testClient(pair.server); 903 defer client.deinit(); 904 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 905 } 906 907 { 908 // wrong connection header 909 var pair = t.SocketPair.init(.{}); 910 defer pair.deinit(); 911 var writer = pair.client.writer(t.io, &.{}); 912 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: something\r\n\r\n"); 913 914 var client = testClient(pair.server); 915 defer client.deinit(); 916 try t.expectError(error.InvalidConnectionHeader, client.handshake("/", .{})); 917 } 918 919 { 920 // missing Sec-Websocket-Accept header 921 var pair = t.SocketPair.init(.{}); 922 defer pair.deinit(); 923 var writer = pair.client.writer(t.io, &.{}); 924 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\nConnection: upgrade\r\n\r\n"); 925 926 var client = testClient(pair.server); 927 defer client.deinit(); 928 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 929 } 930 931 { 932 // wrong Sec-Websocket-Accept header 933 var pair = t.SocketPair.init(.{}); 934 defer pair.deinit(); 935 var writer = pair.client.writer(t.io, &.{}); 936 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: hack\r\n\r\n"); 937 938 var client = testClient(pair.server); 939 defer client.deinit(); 940 try t.expectError(error.InvalidWebsocketAcceptHeader, client.handshake("/", .{})); 941 } 942 943 { 944 // ok for successful 945 var pair = t.SocketPair.init(.{}); 946 defer pair.deinit(); 947 var writer = pair.client.writer(t.io, &.{}); 948 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: C/0nmHhBztSRGR1CwL6Tf4ZjwpY=\r\n\r\n"); 949 950 var client = testClient(pair.server); 951 defer client.deinit(); 952 try client.handshake("/", .{}); 953 try t.expectEqual(0, client._reader.pos); 954 } 955 956 { 957 // ok for successful, with overread 958 var pair = t.SocketPair.init(.{}); 959 defer pair.deinit(); 960 var writer = pair.client.writer(t.io, &.{}); 961 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: C/0nmHhBztSRGR1CwL6Tf4ZjwpY=\r\n\r\nSome Random Data Which is Part Of the Next Message"); 962 963 var client = testClient(pair.server); 964 defer client.deinit(); 965 try client.handshake("/", .{}); 966 try t.expectEqual(50, client._reader.pos); 967 } 968} 969 970test "Client: write/read" { 971 var client = try Client.init(t.io, t.allocator, .{ 972 .port = 9292, 973 .host = "127.0.0.1", 974 }); 975 defer client.deinit(); 976 977 try client.handshake("/", .{ 978 .timeout_ms = 1000, 979 }); 980 981 var buf = [_]u8{ 'o', 'v', 'e', 'r' }; 982 try client.write(&buf); 983 try client.readTimeout(1000); 984 985 const message = (try client.read()) orelse unreachable; 986 try t.expectEqual(.text, message.type); 987 try t.expectString("9000", message.data); 988 989 client.close(.{}) catch unreachable; 990} 991 992test "Client: close with code" { 993 var client = try Client.init(t.io, t.allocator, .{ 994 .port = 9292, 995 .host = "127.0.0.1", 996 }); 997 defer client.deinit(); 998 999 try client.handshake("/", .{ 1000 .timeout_ms = 1000, 1001 }); 1002 1003 client.close(.{ .code = 4002 }) catch unreachable; 1004} 1005 1006test "Client: with code and reason" { 1007 var client = try Client.init(t.io, t.allocator, .{ 1008 .port = 9292, 1009 .host = "127.0.0.1", 1010 }); 1011 defer client.deinit(); 1012 1013 try client.handshake("/", .{ 1014 .timeout_ms = 1000, 1015 }); 1016 1017 client.close(.{ .code = 4002, .reason = "goodbye" }) catch unreachable; 1018} 1019 1020test "Client: Handler" { 1021 var h = try ClientHandler.init(t.io, t.allocator); 1022 defer h.deinit(); 1023 1024 var buf: [6]u8 = undefined; 1025 { 1026 @memcpy(buf[0..3], "dyn"); 1027 try h.client.write(buf[0..3]); 1028 } 1029 1030 { 1031 @memcpy(buf[0..4], "ping"); 1032 try h.client.write(buf[0..4]); 1033 } 1034 1035 { 1036 @memcpy(buf[0..4], "pong"); 1037 try h.client.write(buf[0..4]); 1038 } 1039 1040 { 1041 @memcpy(buf[0..6], "close1"); 1042 try h.client.write(buf[0..6]); 1043 } 1044 1045 try h.client.readLoop(&h); 1046 1047 // if pong is true then ping and message have to be true 1048 // because each asserts the previous 1049 try t.expectEqual(true, h.pong); 1050 try t.expectEqual(true, h.closed); 1051} 1052 1053fn testClient(stream: Io.net.Stream) Client { 1054 const bp = t.allocator.create(buffer.Provider) catch unreachable; 1055 bp.* = buffer.Provider.init(t.io, t.allocator, .{ .count = 0, .size = 0, .max = 4096 }) catch unreachable; 1056 1057 const reader_buf = bp.allocator.alloc(u8, 1024) catch unreachable; 1058 1059 return .{ 1060 .io = t.io, 1061 ._closed = false, 1062 ._own_bp = true, 1063 ._mask_fn = generateMask, 1064 ._compression_opts = null, 1065 .stream = .{ .io = t.io, .stream = stream }, 1066 ._reader = Reader.init(reader_buf, bp, null), 1067 }; 1068} 1069 1070const ClientHandler = struct { 1071 ping: bool = false, 1072 pong: bool = false, 1073 closed: bool = false, 1074 message: bool = false, 1075 client: Client, 1076 1077 fn init(io: Io, allocator: Allocator) !ClientHandler { 1078 var client = try Client.init(io, allocator, .{ 1079 .port = 9292, 1080 .host = "127.0.0.1", 1081 }); 1082 errdefer client.deinit(); 1083 1084 try client.handshake("/", .{ 1085 .timeout_ms = 1000, 1086 }); 1087 1088 return .{ 1089 .client = client, 1090 }; 1091 } 1092 1093 fn deinit(self: *ClientHandler) void { 1094 self.client.deinit(); 1095 } 1096 1097 pub fn serverMessage(self: *ClientHandler, data: []u8, tpe: proto.Message.TextType) !void { 1098 try t.expectEqual(.text, tpe); 1099 try t.expectString("over 9000!", data); 1100 self.message = true; 1101 } 1102 1103 pub fn serverPing(self: *ClientHandler, data: []u8) !void { 1104 try t.expectEqual(true, self.message); 1105 try t.expectString("a-ping", data); 1106 self.ping = true; 1107 } 1108 1109 pub fn serverPong(self: *ClientHandler, data: []u8) !void { 1110 try t.expectEqual(true, self.ping); 1111 try t.expectString("a-pong", data); 1112 self.pong = true; 1113 } 1114 1115 pub fn close(self: *ClientHandler) void { 1116 self.client.close(.{}) catch unreachable; 1117 self.closed = true; 1118 } 1119};