forked from
zzstoatzz.io/websocket.zig
websocket
3.7 kB
124 lines
1const std = @import("std");
2const t = @import("t.zig");
3const ws = @import("websocket.zig");
4
5pub fn init() Testing {
6 return Testing.init();
7}
8
9pub const Testing = struct {
10 closed: bool,
11 conn: ws.Conn,
12 pair: t.SocketPair,
13 reader: ws.proto.Reader,
14 arena: *std.heap.ArenaAllocator,
15
16 received: std.ArrayList(ws.Message),
17 received_index: usize,
18
19 const Opts = struct {
20 port: ?u16 = null,
21 };
22 fn init(opts: Opts) Testing {
23 const arena = t.allocator.create(std.heap.ArenaAllocator) catch unreachable;
24 errdefer t.allocator.destroy(arena);
25
26 arena.* = std.heap.ArenaAllocator.init(t.allocator);
27 errdefer arena.deinit();
28
29 const port = opts.port orelse 0;
30 const pair = t.SocketPair.init(.{ .port = port });
31 const timeout = std.mem.toBytes(std.posix.timeval{
32 .sec = 0,
33 .usec = 50_000,
34 });
35 std.posix.setsockopt(pair.client.socket.handle, std.posix.SOL.SOCKET, std.posix.SO.RCVTIMEO, &timeout) catch unreachable;
36
37 const aa = arena.allocator();
38 const buffer_provider = aa.create(ws.buffer.Provider) catch unreachable;
39 buffer_provider.* = ws.buffer.Provider.init(aa, .{
40 .size = 0,
41 .count = 0,
42 .max = 20_971_520,
43 }) catch unreachable;
44
45 const reader_buf = aa.alloc(u8, 1024) catch unreachable;
46 const reader = ws.proto.Reader.init(reader_buf, buffer_provider, null);
47
48 return .{
49 .closed = false,
50 .pair = pair,
51 .arena = arena,
52 .conn = .{
53 ._closed = false,
54 .started = 0,
55 .stream = pair.server,
56 .address = std.Io.net.IpAddress.parse("127.0.0.1", port) catch unreachable,
57 },
58 .reader = reader,
59 .received = .empty,
60 .received_index = 0,
61 };
62 }
63
64 pub fn deinit(self: *Testing) void {
65 self.pair.deinit();
66 self.arena.deinit();
67 t.allocator.destroy(self.arena);
68 }
69
70 pub fn expectMessage(self: *Testing, op: ws.Message.Type, data: []const u8) !void {
71 try self.ensureMessage();
72
73 const message = self.received.items[self.received_index];
74 self.received_index += 1;
75
76 try t.expectEqual(op, message.type);
77 if (op == .text) {
78 try t.expectString(data, message.data);
79 } else {
80 try t.expectSlice(u8, data, message.data);
81 }
82 }
83
84 pub fn expectClose(self: *Testing) !void {
85 if (self.closed) {
86 return;
87 }
88
89 self.fill() catch if (self.closed) {
90 return;
91 };
92
93 return error.NotClosed;
94 }
95
96 // we have a 50ms timeout on this socket. It's all localhost. We expect
97 // to be able to read messages in that time.
98 pub fn ensureMessage(self: *Testing) !void {
99 if (self.received_index < self.received.items.len) {
100 return;
101 }
102 return self.fill();
103 }
104
105 fn fill(self: *Testing) !void {
106 var client_reader = self.pair.clientReader();
107 self.reader.fill(&client_reader) catch |err| switch (err) {
108 error.WouldBlock => return error.NoMoreData,
109 else => {
110 self.closed = true;
111 return err;
112 },
113 };
114
115 while (true) {
116 const more, const message = (try self.reader.read()) orelse return error.NoMoreData;
117 const aa = self.arena.allocator();
118 self.received.append(aa, message) catch unreachable;
119 if (more == false) {
120 return;
121 }
122 }
123 }
124};