A websocket implementation for 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};