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 1112 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. The caller drains its buffer before reaching here, so polling 452 // only gates an actual socket read. 453 if (self.read_timeout_ms > 0) { 454 var pfd = [_]std.posix.pollfd{.{ 455 .fd = self.stream.socket.handle, 456 .events = std.posix.POLL.IN, 457 .revents = 0, 458 }}; 459 const ready = std.posix.poll(&pfd, @intCast(self.read_timeout_ms)) catch 0; 460 if (ready == 0) return error.WouldBlock; 461 } 462 if (self.tls_client) |tls_client| { 463 var w: std.Io.Writer = .fixed(buf); 464 while (true) { 465 const n = try tls_client.client.reader.stream(&w, .limited(buf.len)); 466 if (n != 0) { 467 return n; 468 } 469 } 470 } 471 return posix.read(self.stream.socket.handle, buf); 472 } 473 474 pub fn writeAll(self: *Stream, data: []const u8) !void { 475 if (self.tls_client) |tls_client| { 476 try tls_client.client.writer.writeAll(data); 477 // I know this looks silly, but as far as I can tell, this is what 478 // we need to do. 479 try tls_client.client.writer.flush(); 480 try tls_client.stream_writer.interface.flush(); 481 return; 482 } 483 484 var writer = self.stream.writer(self.io, &.{}); 485 try writer.interface.writeAll(data); 486 return writer.interface.flush(); 487 } 488 489 const zero_timeout = std.mem.toBytes(posix.timeval{ .sec = 0, .usec = 0 }); 490 pub fn writeTimeout(self: *const Stream, ms: u32) !void { 491 return self.setTimeout(posix.SO.SNDTIMEO, ms); 492 } 493 494 pub fn readTimeout(self: *Stream, ms: u32) !void { 495 // Stored and applied via poll() in read(); see the note there for why this 496 // does not use SO_RCVTIMEO. 497 self.read_timeout_ms = ms; 498 } 499 500 fn setTimeout(self: *const Stream, opt_name: u32, ms: u32) !void { 501 if (ms == 0) { 502 return self.setsockopt(opt_name, &zero_timeout); 503 } 504 505 const timeout = std.mem.toBytes(posix.timeval{ 506 .sec = @intCast(@divTrunc(ms, 1000)), 507 .usec = @intCast(@mod(ms, 1000) * 1000), 508 }); 509 return self.setsockopt(opt_name, &timeout); 510 } 511 512 pub fn setsockopt(self: *const Stream, opt_name: u32, value: []const u8) !void { 513 return posix.setsockopt(self.stream.socket.handle, posix.SOL.SOCKET, opt_name, value); 514 } 515}; 516 517const TLSClient = struct { 518 io: Io, 519 client: tls.Client, 520 stream: Io.net.Stream, 521 stream_writer: Io.net.Stream.Writer, 522 stream_reader: Io.net.Stream.Reader, 523 arena: std.heap.ArenaAllocator, 524 525 fn init(io: Io, allocator: Allocator, stream: Io.net.Stream, config: *const Client.Config) !*TLSClient { 526 var arena = std.heap.ArenaAllocator.init(allocator); 527 errdefer arena.deinit(); 528 529 const aa = arena.allocator(); 530 531 // 0.16: Bundle is heap-allocated so we can pass a pointer to TLS 532 // Options.ca.bundle. A single-threaded RwLock is fine here because 533 // the bundle is only touched by this TLS client; the RwLock serves 534 // only to match the Options.ca.bundle contract. 535 const bundle_ptr = try aa.create(Bundle); 536 if (config.ca_bundle) |existing| { 537 bundle_ptr.* = existing; 538 } else { 539 bundle_ptr.* = .empty; 540 // 0.16: rescan signature is (*Bundle, gpa, io, now: Io.Timestamp). 541 try bundle_ptr.rescan(aa, io, Io.Timestamp.now(io, .real)); 542 } 543 const bundle_lock = try aa.create(Io.RwLock); 544 bundle_lock.* = .init; 545 546 // The TLS input and output have to be max_ciphertext_record_len each. 547 // It isn't clear to me how big the un-encrypted reader and writer 548 // need to be. I would think 0, but that will fail an assertion. I 549 // don't think that it's right that we need 4 buffers, but apparently 550 // we do. Until i figure this out, using 4 x max_ciphertext_record_len 551 // seems like the only safe choice. 552 const buf_len = std.crypto.tls.max_ciphertext_record_len; 553 var buf = try aa.alloc(u8, buf_len * 4); 554 555 const self = try aa.create(TLSClient); 556 self.* = .{ 557 .io = io, 558 .stream = stream, 559 .arena = arena, 560 .client = undefined, 561 .stream_writer = stream.writer(io, buf.ptr[0..buf_len][0..buf_len]), 562 .stream_reader = stream.reader(io, buf.ptr[buf_len .. 2 * buf_len][0..buf_len]), 563 }; 564 565 // 0.16 TLS Client.Options requires `entropy` and `realtime_now` in 566 // addition to the 0.15 set. Fill both from the shim Io — the 567 // entropy buffer is read only during `init`. 568 var entropy_buf: [tls.Client.Options.entropy_len]u8 = undefined; 569 io.random(&entropy_buf); 570 571 self.client = try tls.Client.init( 572 &self.stream_reader.interface, 573 &self.stream_writer.interface, 574 .{ 575 .ca = .{ .bundle = .{ 576 .gpa = aa, 577 .io = io, 578 .lock = bundle_lock, 579 .bundle = bundle_ptr, 580 } }, 581 .host = .{ .explicit = config.host }, 582 .read_buffer = buf.ptr[2 * buf_len .. 3 * buf_len][0..buf_len], 583 .write_buffer = buf.ptr[3 * buf_len .. 4 * buf_len][0..buf_len], 584 .entropy = &entropy_buf, 585 .realtime_now = std.Io.Timestamp.now(io, .real), 586 }, 587 ); 588 589 return self; 590 } 591 592 fn deinit(self: *TLSClient) void { 593 _ = self.client.end() catch {}; 594 self.arena.deinit(); 595 } 596}; 597 598fn generateKey(io: Io) [16]u8 { 599 if (comptime @import("builtin").is_test) { 600 return [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; 601 } 602 var key: [16]u8 = undefined; 603 io.random(&key); 604 return key; 605} 606 607fn generateMask(io: Io) [4]u8 { 608 var m: [4]u8 = undefined; 609 io.random(&m); 610 return m; 611} 612 613fn sendHandshake(path: []const u8, key: []const u8, buf: []u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !void { 614 @memcpy(buf[0..4], "GET "); 615 var pos: usize = 4; 616 var end = pos + path.len; 617 618 { 619 @memcpy(buf[pos..end], path); 620 pos = end; 621 } 622 623 { 624 const headers = " HTTP/1.1\r\ncontent-length: 0\r\nupgrade: websocket\r\nsec-websocket-version: 13\r\nconnection: upgrade\r\nsec-websocket-key: "; 625 end = pos + headers.len; 626 @memcpy(buf[pos..end], headers); 627 628 pos = end; 629 end = pos + key.len; 630 @memcpy(buf[pos..end], key); 631 } 632 633 if (compression) { 634 // NOTE: client_max_window_bits is unsupported 635 const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover"; 636 pos = end; 637 end = pos + permessage_deflate.len; 638 @memcpy(buf[pos..end], permessage_deflate); 639 } 640 641 { 642 pos = end; 643 end = pos + 2; 644 @memcpy(buf[pos..end], "\r\n"); 645 pos = end; 646 } 647 648 if (opts.headers) |extra_headers| { 649 end = pos + extra_headers.len; 650 @memcpy(buf[pos..end], extra_headers); 651 pos = end; 652 if (!std.mem.endsWith(u8, extra_headers, "\r\n")) { 653 buf[pos] = '\r'; 654 buf[pos + 1] = '\n'; 655 pos += 2; 656 } 657 } 658 buf[pos] = '\r'; 659 buf[pos + 1] = '\n'; 660 661 try stream.writeTimeout(opts.timeout_ms); 662 try stream.writeAll(buf[0 .. pos + 2]); 663 try stream.writeTimeout(0); 664} 665 666const HandShakeReply = struct { 667 compression: bool, 668 over_read: usize, 669 670 fn read(io: Io, buf: []u8, key: []const u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !HandShakeReply { 671 const timeout_ms = opts.timeout_ms; 672 // 0.16 removed `std.time.milliTimestamp`; compute ms since epoch 673 // from `std.Io.Timestamp.now(io, .real)` (nanoseconds). 674 const deadline = @divTrunc(Io.Timestamp.now(io, .real).nanoseconds, std.time.ns_per_ms) + timeout_ms; 675 try stream.readTimeout(timeout_ms); 676 677 var pos: usize = 0; 678 var line_start: usize = 0; 679 var complete_response: u8 = 0; 680 var server_compression: bool = false; 681 682 while (true) { 683 const n = stream.read(buf[pos..]) catch |err| { 684 // `error.WouldBlock` may not be in `err`'s set on Windows 685 // (where the read goes through ReadFile), so match by name. 686 if (std.mem.eql(u8, @errorName(err), "WouldBlock")) return error.Timeout; 687 return err; 688 }; 689 if (n == 0) { 690 return error.ConnectionClosed; 691 } 692 693 pos += n; 694 while (std.mem.indexOfScalar(u8, buf[line_start..pos], '\r')) |relative_end| { 695 if (relative_end == 0) { 696 if (complete_response != 15) { 697 return error.InvalidHandshakeResponse; 698 } 699 const over_read = pos - (line_start + 2); 700 std.mem.copyForwards(u8, buf[0..over_read], buf[line_start + 2 .. pos]); 701 try stream.readTimeout(0); 702 return .{ 703 .over_read = over_read, 704 .compression = server_compression, 705 }; 706 } 707 708 const line_end = line_start + relative_end; 709 const line = buf[line_start..line_end]; 710 711 // the next line starts where this line ends, skip over the \r\n 712 line_start = line_end + 2; 713 714 if (complete_response == 0) { 715 if (!ascii.startsWithIgnoreCase(line, "HTTP/1.1 101 ")) { 716 return error.InvalidHandshakeResponse; 717 } 718 complete_response |= 1; 719 continue; 720 } 721 722 for (line, 0..) |b, i| { 723 // find the colon and lowercase the header while we're iterating 724 if ('A' <= b and b <= 'Z') { 725 line[i] = b + 32; 726 continue; 727 } 728 729 if (b != ':') { 730 continue; 731 } 732 733 switch (i) { 734 7 => if (std.mem.eql(u8, line[0..i], "upgrade")) { 735 if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "websocket")) { 736 return error.InvalidUpgradeHeader; 737 } 738 complete_response |= 2; 739 }, 740 10 => if (std.mem.eql(u8, line[0..i], "connection")) { 741 if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "upgrade")) { 742 return error.InvalidConnectionHeader; 743 } 744 complete_response |= 4; 745 }, 746 20 => if (std.mem.eql(u8, line[0..i], "sec-websocket-accept")) { 747 var h: [20]u8 = undefined; 748 { 749 var hasher = std.crypto.hash.Sha1.init(.{}); 750 hasher.update(key); 751 hasher.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); 752 hasher.final(&h); 753 } 754 755 var encoded_buf: [28]u8 = undefined; 756 const sec_hash = std.base64.standard.Encoder.encode(&encoded_buf, &h); 757 const header_value = std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace); 758 759 if (!std.mem.eql(u8, header_value, sec_hash)) { 760 return error.InvalidWebsocketAcceptHeader; 761 } 762 complete_response |= 8; 763 }, 764 24 => if (std.mem.eql(u8, line[0..i], "sec-websocket-extensions")) { 765 if (try parseExtension(line[i + 1 ..])) |sc| { 766 if (!compression) { 767 // server is saying compression, but we didn't ask for it. 768 return error.InvalidExtensionHeader; 769 } 770 if (!sc.client_no_context_takeover or !sc.server_no_context_takeover) { 771 // as of Zig 0.15, we no longer support context takeover 772 // We told the server this, it should have respected it. 773 return error.InvalidExtensionHeader; 774 } 775 776 server_compression = true; 777 } 778 }, 779 else => {}, // some other header we don't care about 780 } 781 } 782 } 783 784 if (@divTrunc(Io.Timestamp.now(io, .real).nanoseconds, std.time.ns_per_ms) > deadline) { 785 return error.Timeout; 786 } 787 788 if (pos == buf.len) { 789 return error.ResponseTooLarge; 790 } 791 } 792 } 793 794 pub fn parseExtension(value: []const u8) !?ServerHandshake.Compression { 795 var deflate = false; 796 var client_max_bits: u8 = 15; 797 var client_no_context_takeover = false; 798 var server_no_context_takeover = false; 799 800 var it = std.mem.splitScalar(u8, value, ';'); 801 while (it.next()) |param_| { 802 const param = std.mem.trim(u8, param_, &ascii.whitespace); 803 if (std.mem.eql(u8, param, "permessage-deflate")) { 804 deflate = true; 805 continue; 806 } 807 if (std.mem.eql(u8, param, "client_no_context_takeover")) { 808 client_no_context_takeover = true; 809 continue; 810 } 811 if (std.mem.eql(u8, param, "server_no_context_takeover")) { 812 server_no_context_takeover = true; 813 continue; 814 } 815 const client_max_window_bits = "client_max_window_bits="; 816 if (std.mem.startsWith(u8, param, client_max_window_bits)) { 817 client_max_bits = std.fmt.parseInt(u8, param[client_max_window_bits.len..], 10) catch { 818 return error.InvalidCompressionServerMaxBits; 819 }; 820 } 821 } 822 if (deflate == false) { 823 return null; 824 } 825 826 if (client_max_bits != 15) { 827 // We don't offer client window, so if the server asks for one, that's an error 828 return error.InvalidExtensionHeader; 829 } 830 831 return .{ 832 .client_no_context_takeover = client_no_context_takeover, 833 .server_no_context_takeover = server_no_context_takeover, 834 }; 835 } 836}; 837 838const t = @import("../t.zig"); 839test "Client: handshake" { 840 { 841 // empty response 842 var pair = t.SocketPair.init(.{}); 843 defer pair.deinit(); 844 var writer = pair.client.writer(t.io, &.{}); 845 try writer.interface.writeAll("\r\n\r\n"); 846 847 var client = testClient(pair.server); 848 defer client.deinit(); 849 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 850 } 851 852 { 853 // invalid websocket response 854 var pair = t.SocketPair.init(.{}); 855 defer pair.deinit(); 856 var writer = pair.client.writer(t.io, &.{}); 857 try writer.interface.writeAll("HTTP/1.1 200 OK\r\n\r\n"); 858 859 var client = testClient(pair.server); 860 defer client.deinit(); 861 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 862 } 863 864 { 865 // missing upgrade header 866 var pair = t.SocketPair.init(.{}); 867 defer pair.deinit(); 868 var writer = pair.client.writer(t.io, &.{}); 869 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\n\r\n"); 870 871 var client = testClient(pair.server); 872 defer client.deinit(); 873 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 874 } 875 876 { 877 // wrong upgrade header 878 var pair = t.SocketPair.init(.{}); 879 defer pair.deinit(); 880 var writer = pair.client.writer(t.io, &.{}); 881 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: nope\r\n\r\n"); 882 883 var client = testClient(pair.server); 884 defer client.deinit(); 885 try t.expectError(error.InvalidUpgradeHeader, client.handshake("/", .{})); 886 } 887 888 { 889 // missing connection header 890 var pair = t.SocketPair.init(.{}); 891 defer pair.deinit(); 892 var writer = pair.client.writer(t.io, &.{}); 893 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\n\r\n"); 894 895 var client = testClient(pair.server); 896 defer client.deinit(); 897 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 898 } 899 900 { 901 // wrong connection header 902 var pair = t.SocketPair.init(.{}); 903 defer pair.deinit(); 904 var writer = pair.client.writer(t.io, &.{}); 905 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: something\r\n\r\n"); 906 907 var client = testClient(pair.server); 908 defer client.deinit(); 909 try t.expectError(error.InvalidConnectionHeader, client.handshake("/", .{})); 910 } 911 912 { 913 // missing Sec-Websocket-Accept header 914 var pair = t.SocketPair.init(.{}); 915 defer pair.deinit(); 916 var writer = pair.client.writer(t.io, &.{}); 917 try writer.interface.writeAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\nConnection: upgrade\r\n\r\n"); 918 919 var client = testClient(pair.server); 920 defer client.deinit(); 921 try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); 922 } 923 924 { 925 // wrong Sec-Websocket-Accept header 926 var pair = t.SocketPair.init(.{}); 927 defer pair.deinit(); 928 var writer = pair.client.writer(t.io, &.{}); 929 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"); 930 931 var client = testClient(pair.server); 932 defer client.deinit(); 933 try t.expectError(error.InvalidWebsocketAcceptHeader, client.handshake("/", .{})); 934 } 935 936 { 937 // ok for successful 938 var pair = t.SocketPair.init(.{}); 939 defer pair.deinit(); 940 var writer = pair.client.writer(t.io, &.{}); 941 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"); 942 943 var client = testClient(pair.server); 944 defer client.deinit(); 945 try client.handshake("/", .{}); 946 try t.expectEqual(0, client._reader.pos); 947 } 948 949 { 950 // ok for successful, with overread 951 var pair = t.SocketPair.init(.{}); 952 defer pair.deinit(); 953 var writer = pair.client.writer(t.io, &.{}); 954 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"); 955 956 var client = testClient(pair.server); 957 defer client.deinit(); 958 try client.handshake("/", .{}); 959 try t.expectEqual(50, client._reader.pos); 960 } 961} 962 963test "Client: write/read" { 964 var client = try Client.init(t.io, t.allocator, .{ 965 .port = 9292, 966 .host = "127.0.0.1", 967 }); 968 defer client.deinit(); 969 970 try client.handshake("/", .{ 971 .timeout_ms = 1000, 972 }); 973 974 var buf = [_]u8{ 'o', 'v', 'e', 'r' }; 975 try client.write(&buf); 976 try client.readTimeout(1000); 977 978 const message = (try client.read()) orelse unreachable; 979 try t.expectEqual(.text, message.type); 980 try t.expectString("9000", message.data); 981 982 client.close(.{}) catch unreachable; 983} 984 985test "Client: close with code" { 986 var client = try Client.init(t.io, t.allocator, .{ 987 .port = 9292, 988 .host = "127.0.0.1", 989 }); 990 defer client.deinit(); 991 992 try client.handshake("/", .{ 993 .timeout_ms = 1000, 994 }); 995 996 client.close(.{ .code = 4002 }) catch unreachable; 997} 998 999test "Client: with code and reason" { 1000 var client = try Client.init(t.io, t.allocator, .{ 1001 .port = 9292, 1002 .host = "127.0.0.1", 1003 }); 1004 defer client.deinit(); 1005 1006 try client.handshake("/", .{ 1007 .timeout_ms = 1000, 1008 }); 1009 1010 client.close(.{ .code = 4002, .reason = "goodbye" }) catch unreachable; 1011} 1012 1013test "Client: Handler" { 1014 var h = try ClientHandler.init(t.io, t.allocator); 1015 defer h.deinit(); 1016 1017 var buf: [6]u8 = undefined; 1018 { 1019 @memcpy(buf[0..3], "dyn"); 1020 try h.client.write(buf[0..3]); 1021 } 1022 1023 { 1024 @memcpy(buf[0..4], "ping"); 1025 try h.client.write(buf[0..4]); 1026 } 1027 1028 { 1029 @memcpy(buf[0..4], "pong"); 1030 try h.client.write(buf[0..4]); 1031 } 1032 1033 { 1034 @memcpy(buf[0..6], "close1"); 1035 try h.client.write(buf[0..6]); 1036 } 1037 1038 try h.client.readLoop(&h); 1039 1040 // if pong is true then ping and message have to be true 1041 // because each asserts the previous 1042 try t.expectEqual(true, h.pong); 1043 try t.expectEqual(true, h.closed); 1044} 1045 1046fn testClient(stream: Io.net.Stream) Client { 1047 const bp = t.allocator.create(buffer.Provider) catch unreachable; 1048 bp.* = buffer.Provider.init(t.io, t.allocator, .{ .count = 0, .size = 0, .max = 4096 }) catch unreachable; 1049 1050 const reader_buf = bp.allocator.alloc(u8, 1024) catch unreachable; 1051 1052 return .{ 1053 .io = t.io, 1054 ._closed = false, 1055 ._own_bp = true, 1056 ._mask_fn = generateMask, 1057 ._compression_opts = null, 1058 .stream = .{ .io = t.io, .stream = stream }, 1059 ._reader = Reader.init(reader_buf, bp, null), 1060 }; 1061} 1062 1063const ClientHandler = struct { 1064 ping: bool = false, 1065 pong: bool = false, 1066 closed: bool = false, 1067 message: bool = false, 1068 client: Client, 1069 1070 fn init(io: Io, allocator: Allocator) !ClientHandler { 1071 var client = try Client.init(io, allocator, .{ 1072 .port = 9292, 1073 .host = "127.0.0.1", 1074 }); 1075 errdefer client.deinit(); 1076 1077 try client.handshake("/", .{ 1078 .timeout_ms = 1000, 1079 }); 1080 1081 return .{ 1082 .client = client, 1083 }; 1084 } 1085 1086 fn deinit(self: *ClientHandler) void { 1087 self.client.deinit(); 1088 } 1089 1090 pub fn serverMessage(self: *ClientHandler, data: []u8, tpe: proto.Message.TextType) !void { 1091 try t.expectEqual(.text, tpe); 1092 try t.expectString("over 9000!", data); 1093 self.message = true; 1094 } 1095 1096 pub fn serverPing(self: *ClientHandler, data: []u8) !void { 1097 try t.expectEqual(true, self.message); 1098 try t.expectString("a-ping", data); 1099 self.ping = true; 1100 } 1101 1102 pub fn serverPong(self: *ClientHandler, data: []u8) !void { 1103 try t.expectEqual(true, self.ping); 1104 try t.expectString("a-pong", data); 1105 self.pong = true; 1106 } 1107 1108 pub fn close(self: *ClientHandler) void { 1109 self.client.close(.{}) catch unreachable; 1110 self.closed = true; 1111 } 1112};