const 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 }; } };