Something went wrong. Try again.
websocket
Something went wrong. Try again.
7.6 kB · 248 lines
Zig
at main
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249const std = @import("std");const proto = @import("proto.zig");
const Io = std.Io;const net = Io.net;const ArrayList = std.ArrayList;
const Message = proto.Message;
pub const allocator = std.testing.allocator;
pub fn expectEqual(expected: anytype, actual: anytype) !void { try std.testing.expectEqual(expected, actual);}
pub const expectError = std.testing.expectError;pub const expectString = std.testing.expectEqualStrings;pub const expectSlice = std.testing.expectEqualSlices;
pub fn getRandom() std.Random.DefaultPrng { const io = std.Options.debug_io; var seed_bytes: [8]u8 = undefined; io.random(&seed_bytes); const seed: u64 = @bitCast(seed_bytes); return std.Random.DefaultPrng.init(seed);}
pub var arena = std.heap.ArenaAllocator.init(allocator);pub fn reset() void { _ = arena.reset(.free_all);}
pub const Writer = struct { pos: usize, buf: std.ArrayList(u8), random: std.Random.DefaultPrng,
pub fn init() Writer { return .{ .pos = 0, .buf = .empty, .random = getRandom(), }; }
pub fn deinit(self: *Writer) void { self.buf.deinit(allocator); }
pub fn ping(self: *Writer) void { return self.pingPayload(""); }
pub fn pong(self: *Writer) void { return self.frame(true, 10, "", 0); }
pub fn pingPayload(self: *Writer, payload: []const u8) void { return self.frame(true, 9, payload, 0); }
pub fn textFrame(self: *Writer, fin: bool, payload: []const u8) void { return self.frame(fin, 1, payload, 0); }
pub fn cont(self: *Writer, fin: bool, payload: []const u8) void { return self.frame(fin, 0, payload, 0); }
pub fn frame(self: *Writer, fin: bool, op_code: u8, payload: []const u8, reserved: u8) void { var buf = &self.buf;
const l = payload.len; var length_of_length: usize = 0;
if (l > 125) { if (l < 65536) { length_of_length = 2; } else { length_of_length = 8; } }
// 2 byte header + length_of_length + mask + payload_length const needed = 2 + length_of_length + 4 + l; buf.ensureUnusedCapacity(allocator, needed) catch unreachable;
if (fin) { buf.appendAssumeCapacity(128 | op_code | reserved); } else { buf.appendAssumeCapacity(op_code | reserved); }
if (length_of_length == 0) { buf.appendAssumeCapacity(128 | @as(u8, @intCast(l))); } else if (length_of_length == 2) { buf.appendAssumeCapacity(128 | 126); buf.appendAssumeCapacity(@intCast((l >> 8) & 0xFF)); buf.appendAssumeCapacity(@intCast(l & 0xFF)); } else { buf.appendAssumeCapacity(128 | 127); buf.appendAssumeCapacity(@intCast((l >> 56) & 0xFF)); buf.appendAssumeCapacity(@intCast((l >> 48) & 0xFF)); buf.appendAssumeCapacity(@intCast((l >> 40) & 0xFF)); buf.appendAssumeCapacity(@intCast((l >> 32) & 0xFF)); buf.appendAssumeCapacity(@intCast((l >> 24) & 0xFF)); buf.appendAssumeCapacity(@intCast((l >> 16) & 0xFF)); buf.appendAssumeCapacity(@intCast((l >> 8) & 0xFF)); buf.appendAssumeCapacity(@intCast(l & 0xFF)); }
var mask: [4]u8 = undefined; self.random.random().bytes(&mask);
buf.appendSliceAssumeCapacity(&mask); for (payload, 0..) |b, i| { buf.appendAssumeCapacity(b ^ mask[i & 3]); } }
pub fn bytes(self: *const Writer) []const u8 { return self.buf.items; }
pub fn clear(self: *Writer) void { self.pos = 0; self.buf.clearRetainingCapacity(); }
pub fn read( self: *Writer, buf: []u8, ) !usize { const data = self.buf.items[self.pos..];
if (data.len == 0 or buf.len == 0) { return 0; }
// randomly fragment the data const to_read = self.random.random().intRangeAtMost(usize, 1, @min(data.len, buf.len)); @memcpy(buf[0..to_read], data[0..to_read]); self.pos += to_read; return to_read; }};
/// Reads directly from a socket, preserving SO_RCVTIMEO as WouldBlock.pub const StreamReader = struct { handle: net.Socket.Handle,
pub fn read(self: *StreamReader, buf: []u8) !usize { return std.posix.read(self.handle, buf) catch |err| { return switch (err) { error.ConnectionResetByPeer => error.ConnectionResetByPeer, error.WouldBlock => error.WouldBlock, else => error.Unexpected, }; }; }};
pub const SocketPair = struct { writer: Writer, io: Io, client: net.Stream, server: net.Stream, server_taken: bool = false,
const Opts = struct { port: ?u16 = null, };
pub fn init(opts: Opts) SocketPair { const io = std.Options.debug_io; const port: u16 = opts.port orelse 0;
// use Io.net to create a listener, connect a client, and accept var listen_addr: net.IpAddress = .{ .ip4 = .loopback(port) }; var server = net.IpAddress.listen(&listen_addr, io, .{ .reuse_address = true }) catch unreachable;
// get the actual bound address (for ephemeral port) const bound_addr = server.socket.address;
// connect client const client = net.IpAddress.connect(&bound_addr, io, .{ .mode = .stream }) catch unreachable;
// accept server-side connection const accepted = server.accept(io) catch unreachable; server.deinit(io);
return .{ .io = io, .client = client, .server = accepted, .writer = Writer.init(), }; }
pub fn deinit(self: *SocketPair) void { self.writer.deinit(); self.client.close(self.io); if (!self.server_taken) self.server.close(self.io); }
pub fn pingPayload(self: *SocketPair, payload: []const u8) void { self.writer.pingPayload(payload); }
pub fn textFrame(self: *SocketPair, fin: bool, payload: []const u8) void { self.writer.textFrame(fin, payload); }
pub fn cont(self: *SocketPair, fin: bool, payload: []const u8) void { self.writer.cont(fin, payload); }
pub fn sendBuf(self: *SocketPair) void { self.ioWriteAll(self.writer.bytes()) catch unreachable; self.writer.clear(); }
pub fn clientWriteAll(self: *SocketPair, data: []const u8) !void { try self.ioWriteAll(data); }
fn ioWriteAll(self: *SocketPair, data: []const u8) !void { var remaining = data; while (remaining.len > 0) { // netWrite: header is sent first, data array's last element is the splat pattern. // Pass remaining as header, empty pattern with splat=0. const empty = [_][]const u8{""}; const n = self.io.vtable.netWrite(self.io.userdata, self.client.socket.handle, remaining, &empty, 0) catch { return error.SendError; }; if (n == 0) return error.SendError; remaining = remaining[n..]; } }
pub fn serverReader(self: *SocketPair) StreamReader { return .{ .handle = self.server.socket.handle }; }
pub fn clientReader(self: *SocketPair) StreamReader { return .{ .handle = self.client.socket.handle }; }};