forked from
zzstoatzz.io/websocket.zig
websocket
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};