Something went wrong. Try again.
websocket
Something went wrong. Try again.
55 kB · 1443 lines
Zig
at main
1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117111811191120112111221123112411251126112711281129113011311132113311341135113611371138113911401141114211431144114511461147114811491150115111521153115411551156115711581159116011611162116311641165116611671168116911701171117211731174117511761177117811791180118111821183118411851186118711881189119011911192119311941195119611971198119912001201120212031204120512061207120812091210121112121213121412151216121712181219122012211222122312241225122612271228122912301231123212331234123512361237123812391240124112421243124412451246124712481249125012511252125312541255125612571258125912601261126212631264126512661267126812691270127112721273127412751276127712781279128012811282128312841285128612871288128912901291129212931294129512961297129812991300130113021303130413051306130713081309131013111312131313141315131613171318131913201321132213231324132513261327132813291330133113321333133413351336133713381339134013411342134313441345134613471348134913501351135213531354135513561357135813591360136113621363136413651366136713681369137013711372137313741375137613771378137913801381138213831384138513861387138813891390139113921393139413951396139713981399140014011402140314041405140614071408140914101411141214131414141514161417141814191420142114221423142414251426142714281429143014311432143314341435143614371438143914401441144214431444const std = @import("std");const proto = @import("../proto.zig");const buffer = @import("../buffer.zig");
const ascii = std.ascii;const Io = std.Io;const net = Io.net;const posix = std.posix;const tls = std.crypto.tls;const log = std.log.scoped(.websocket);
const Reader = proto.Reader;const Allocator = std.mem.Allocator;const Bundle = std.crypto.Certificate.Bundle;const CompressionOpts = @import("../websocket.zig").Compression;const ServerHandshake = @import("../server/handshake.zig").Handshake;
fn milliTimestamp(io: Io) i64 { const ts = Io.Timestamp.now(io, .real); return @intCast(@divTrunc(ts.nanoseconds, std.time.ns_per_ms));}
fn ReadLoopHandler(comptime T: type) type { const info = @typeInfo(T);
switch (info) { .@"struct" => |struct_info| { if (struct_info.is_tuple) @compileError("readLoop: handler does not support tuples.");
return T; }, .pointer => |ptr_info| { switch (ptr_info.size) { .one => return ReadLoopHandler(ptr_info.child), else => @compileError("readLoop: handler does not support Slice, C and Many pointers."), } }, else => @compileError("readLoop: expected handler to be a struct or pointer to a struct but found '" ++ @tagName(info) ++ "'"), }}
pub const Client = struct { io: Io, stream: Stream, _reader: Reader, _closed: bool, _compression_opts: ?CompressionOpts, _compression: ?Client.Compression = null,
/// populated when handshake() fails with error.InvalidHandshakeResponse /// because the server answered the upgrade with a non-101 HTTP response. /// servers reject upgrades pre-connection with meaningful statuses and /// JSON error envelopes (e.g. jetstream's 400 CursorTooOld); without /// this the caller only sees the bare error and cannot distinguish a /// terminal rejection from a transient one. null when the handshake /// failed for any other reason (malformed response, timeout, transport). handshake_failure: ?HandshakeFailure = null,
// Serializes writes from concurrent tasks (ping loop, auto-pong, close). // Matches server-side Conn.lock pattern. _write_lock: Io.Mutex = .init,
// When creating a client, we can either be given a BufferProvider or create // one ourselves. If we create it ourselves (in init), we "own" it and must // free it on deinit. (The reference to the buffer provider is already in the // reader, no need to hold another reference in the client). _own_bp: bool,
// For advanced cases, a custom masking function can be provided. Masking // is a security feature that only really makes sense in the browser. If you // aren't running websockets in the browser AND you control both the client // and the server, you could get a performance boost by not masking. _mask_fn: *const fn (Io) [4]u8,
pub const Config = struct { port: u16, host: []const u8, tls: bool = false, max_size: usize = 65536, buffer_size: usize = 4096, ca_bundle: ?Bundle = null, mask_fn: ?*const fn (Io) [4]u8 = null, buffer_provider: ?*buffer.Provider = null, compression: ?CompressionOpts = null, };
pub const HandshakeFailure = struct { status: u16, body_len: usize = 0, body_buf: [1024]u8 = undefined,
pub fn body(self: *const HandshakeFailure) []const u8 { return self.body_buf[0..self.body_len]; } };
pub const HandshakeOpts = struct { timeout_ms: u32 = 10000, headers: ?[]const u8 = null, };
const Compression = struct { allocator: Allocator, retain_writer: bool, write_treshold: usize, writer: std.Io.Writer.Allocating, };
pub fn init(io: Io, allocator: Allocator, config: Config) !Client { if (config.compression != null) { log.err("Compression is disabled as part of the 0.15 upgrade. I do hope to re-enable it soon.", .{}); return error.InvalidConfiguraion; }
// 0.16: networking via Io.net.HostName const host_name = try net.HostName.init(config.host); // 0.16: connect requires mode option (stream vs datagram) const net_stream = try host_name.connect(io, config.port, .{ .mode = .stream });
var tls_client: ?*TLSClient = null; if (config.tls) { tls_client = try TLSClient.init(allocator, io, net_stream, &config); } const stream = Stream.init(io, net_stream, tls_client);
var own_bp = false; var buffer_provider: *buffer.Provider = undefined;
// If a buffer_provider is provided, we'll use that. // If it isn't, we need to create one which also means we now "own" it // and we're responsible for cleaning it up if (config.buffer_provider) |shared_bp| { buffer_provider = shared_bp; } else { own_bp = true; buffer_provider = try allocator.create(buffer.Provider); errdefer allocator.destroy(buffer_provider); buffer_provider.* = try buffer.Provider.init(allocator, .{ .size = 0, .count = 0, .max = config.max_size, }); }
errdefer if (own_bp) { buffer_provider.deinit(); allocator.destroy(buffer_provider); };
const reader_buf = try buffer_provider.allocator.alloc(u8, config.buffer_size); errdefer buffer_provider.allocator.free(reader_buf);
return .{ .io = io, .stream = stream, ._closed = false, ._own_bp = own_bp, ._mask_fn = config.mask_fn orelse generateMask, ._compression_opts = null, //TODO: ZIG 0.15 ._reader = Reader.init(reader_buf, buffer_provider, null), }; }
// 0.16: Alternative init that accepts a pre-existing stream. // Supports TLS if config.tls is set (uses config.host for SNI). pub fn initWithStream(io: Io, allocator: Allocator, net_stream: net.Stream, config: Config) !Client { var tls_client: ?*TLSClient = null; if (config.tls) { tls_client = try TLSClient.init(allocator, io, net_stream, &config); } const stream = Stream.init(io, net_stream, tls_client);
var own_bp = false; var buffer_provider: *buffer.Provider = undefined;
if (config.buffer_provider) |shared_bp| { buffer_provider = shared_bp; } else { own_bp = true; buffer_provider = try allocator.create(buffer.Provider); errdefer allocator.destroy(buffer_provider); buffer_provider.* = try buffer.Provider.init(allocator, .{ .size = 0, .count = 0, .max = config.max_size, }); }
errdefer if (own_bp) { buffer_provider.deinit(); allocator.destroy(buffer_provider); };
const reader_buf = try buffer_provider.allocator.alloc(u8, config.buffer_size); errdefer buffer_provider.allocator.free(reader_buf);
return .{ .io = io, .stream = stream, ._closed = false, ._own_bp = own_bp, ._mask_fn = config.mask_fn orelse generateMask, ._compression_opts = null, ._reader = Reader.init(reader_buf, buffer_provider, null), }; }
pub fn deinit(self: *Client) void { self.closeStream();
const larger_buffer_provider = self._reader.large_buffer_provider; const allocator = larger_buffer_provider.allocator; allocator.free(self._reader.static);
self._reader.deinit();
if (self._own_bp) { larger_buffer_provider.deinit(); allocator.destroy(larger_buffer_provider); } }
pub fn handshake(self: *Client, path: []const u8, opts: HandshakeOpts) !void { const stream = &self.stream; errdefer self.closeStream();
// we've already setup our reader, and the reader has a static buffer // we might as well use it! const buf = self._reader.static; const key = blk: { const bin_key = generateKey(self.io); var encoded_key: [24]u8 = undefined; break :blk std.base64.standard.Encoder.encode(&encoded_key, &bin_key); };
try sendHandshake(path, key, buf, &opts, self._compression_opts != null, stream);
self.handshake_failure = null; const res = try HandShakeReply.read(buf, key, &opts, self._compression_opts != null, stream, &self.handshake_failure); errdefer self.close(.{ .code = 1001 }) catch unreachable;
// Set up compression with agreed-on parameters if (res.compression) { try self.setupCompression(); }
// We might have read more than handshake response. If so, readHandshakeReply // has positioned the extra data at the start of the buffer, but we need // to set the length. self._reader.pos = res.over_read; }
fn setupCompression(self: *Client) !void { std.debug.assert(self._compression_opts != null); self._reader.allow_compressed = true;
const allocator = self._reader.large_buffer_provider.allocator; const config = self._compression_opts.?; self._compression = .{ .allocator = allocator, .write_treshold = config.write_threshold.?, .retain_writer = config.retain_write_buffer, .writer = std.Io.Writer.Allocating.init(allocator), }; }
pub fn readLoop(self: *Client, handler: anytype) !void { const Handler = ReadLoopHandler(@TypeOf(handler)); var reader = &self._reader;
defer if (comptime std.meta.hasFn(Handler, "close")) { handler.close(); };
// block until we have data try self.readTimeout(0);
while (true) { const message = self.read() catch |err| switch (err) { error.Closed => return, else => return err, } orelse unreachable;
const message_type = message.type; defer reader.done(message_type);
switch (message_type) { .text, .binary => { switch (comptime @typeInfo(@TypeOf(Handler.serverMessage)).@"fn".params.len) { 2 => try handler.serverMessage(message.data), 3 => try handler.serverMessage(message.data, if (message_type == .text) .text else .binary), else => @compileError(@typeName(Handler) ++ ".serverMessage must accept 2 or 3 parameters"), } }, .ping => if (comptime std.meta.hasFn(Handler, "serverPing")) { try handler.serverPing(message.data); } else { // @constCast is safe because we know message.data points to // reader.buffer.buf, which we own and which can be mutated try self.writeFrame(.pong, @constCast(message.data)); }, .close => { if (comptime std.meta.hasFn(Handler, "serverClose")) { try handler.serverClose(message.data); } else { self.close(.{}) catch unreachable; } return; }, .pong => if (comptime std.meta.hasFn(Handler, "serverPong")) { try handler.serverPong(message.data); }, } } }
pub const HeartbeatConfig = struct { /// ping interval in milliseconds. readTimeout is set to this value. /// when no data arrives within the interval, a ping is sent. interval_ms: u32 = 30_000, /// close connection after this many consecutive intervals with no data or pong. max_failures: u32 = 4, };
pub fn readLoopWithHeartbeat(self: *Client, handler: anytype, heartbeat: HeartbeatConfig) !void { const Handler = ReadLoopHandler(@TypeOf(handler)); var reader = &self._reader;
defer if (comptime std.meta.hasFn(Handler, "close")) { handler.close(); };
try self.readTimeout(heartbeat.interval_ms);
var pending_pings: u32 = 0;
while (true) { const message = self.read() catch |err| switch (err) { error.Closed => return, else => return err, } orelse { // timeout — no data in interval pending_pings += 1; if (pending_pings >= heartbeat.max_failures) { self.close(.{}) catch {}; return error.Closed; } self.writePing(&.{}) catch { self.close(.{}) catch {}; return error.Closed; }; continue; };
// any received frame proves liveness pending_pings = 0;
const message_type = message.type; defer reader.done(message_type);
switch (message_type) { .text, .binary => { switch (comptime @typeInfo(@TypeOf(Handler.serverMessage)).@"fn".params.len) { 2 => try handler.serverMessage(message.data), 3 => try handler.serverMessage(message.data, if (message_type == .text) .text else .binary), else => @compileError(@typeName(Handler) ++ ".serverMessage must accept 2 or 3 parameters"), } }, .ping => if (comptime std.meta.hasFn(Handler, "serverPing")) { try handler.serverPing(message.data); } else { try self.writeFrame(.pong, @constCast(message.data)); }, .close => { if (comptime std.meta.hasFn(Handler, "serverClose")) { try handler.serverClose(message.data); } else { self.close(.{}) catch unreachable; } return; }, .pong => if (comptime std.meta.hasFn(Handler, "serverPong")) { try handler.serverPong(message.data); }, } } }
pub fn read(self: *Client) !?proto.Message { var reader = &self._reader; const stream = &self.stream;
while (true) { // try to read a message from our buffer first, before trying to // get more data from the socket. const has_more, const message = reader.read() catch |err| { self.close(.{ .code = 1002 }) catch unreachable; return err; } orelse { // 0.16: Io vtable error set changed reader.fill(stream) catch |err| { // Check for timeout/would-block type errors if (err == error.Canceled or err == error.WouldBlock) return null; // Check for connection closed errors if (err == error.Closed or err == error.ConnectionResetByPeer or err == error.NotOpenForReading) { @atomicStore(bool, &self._closed, true, .monotonic); return error.Closed; } self.close(.{ .code = 1002 }) catch unreachable; return err; }; continue; };
_ = has_more; return message; } }
pub fn done(self: *Client, message: proto.Message) void { self._reader.done(message.type); }
pub fn readLoopInNewThread(self: *Client, h: anytype) !std.Thread { return std.Thread.spawn(.{}, readLoopOwnedThread, .{ self, h }); }
fn readLoopOwnedThread(self: *Client, h: anytype) void { self.readLoop(h) catch {}; }
pub fn writeTimeout(self: *const Client, ms: u32) !void { return self.stream.writeTimeout(ms); }
pub fn readTimeout(self: *const Client, ms: u32) !void { return self.stream.readTimeout(ms); }
pub fn write(self: *Client, data: []u8) !void { return self.writeFrame(.text, data); }
pub fn writeText(self: *Client, data: []u8) !void { return self.writeFrame(.text, data); }
pub fn writeBin(self: *Client, data: []u8) !void { return self.writeFrame(.binary, data); }
pub fn writePing(self: *Client, data: []u8) !void { return self.writeFrame(.ping, data); }
pub fn writePong(self: *Client, data: []u8) !void { return self.writeFrame(.pong, data); }
const CloseOpts = struct { code: ?u16 = null, reason: []const u8 = "", };
pub fn close(self: *Client, opts: CloseOpts) !void { if (@atomicRmw(bool, &self._closed, .Xchg, true, .monotonic) == true) { // already closed return; }
// Hold the write lock across BOTH the close frame and the stream // teardown. stream.close() frees the TLS client, so a concurrent // writer (a keepalive ping task, say) that was queued on the lock // when close ran must observe error.Closed on acquiring it — not // write through freed TLS state. Releasing the lock between the // close frame and stream.close() left exactly that window, and a // production ping fiber died in it (GPF, 2026-08-27). self._write_lock.lockUncancelable(self.io); defer self._write_lock.unlock(self.io); defer self.stream.close();
const code = opts.code orelse { self.writeFrameLocked(.close, "") catch {}; return; };
const reason = opts.reason; if (reason.len > 123) { return error.ReasonTooLong; }
var buf: [125]u8 = undefined; buf[0] = @intCast((code >> 8) & 0xFF); buf[1] = @intCast(code & 0xFF);
const end = 2 + reason.len; @memcpy(buf[2..end], reason); self.writeFrameLocked(.close, buf[0..end]) catch {}; }
pub fn writeFrame(self: *Client, op_code: proto.OpCode, data: []u8) !void { // Serialize writes — concurrent ping/pong/close must not interleave // frames. The closed check must happen under the lock: close() tears // the stream (and TLS state) down while holding it, so a writer that // slept through a close must fail here rather than touch the stream. self._write_lock.lockUncancelable(self.io); defer self._write_lock.unlock(self.io); if (@atomicLoad(bool, &self._closed, .monotonic)) { return error.Closed; } return self.writeFrameLocked(op_code, data); }
/// the write lock must be held (and _closed checked, except by close /// itself) before calling. fn writeFrameLocked(self: *Client, op_code: proto.OpCode, data: []u8) !void { const payload = data; const compressed = false;
// maximum possible prefix length. op_code + length_type + 8byte length + 4 byte mask var buf: [14]u8 = undefined; const header = proto.writeFrameHeader(&buf, op_code, payload.len, compressed);
const header_len = header.len; const header_end = header.len + 4; // for the mask
buf[1] |= 128; // indicate that the payload is masked
const mask = self._mask_fn(self.io); @memcpy(buf[header_len..header_end], &mask);
if (payload.len > 0) { proto.mask(&mask, payload); }
try self.stream.writeAll(buf[0..header_end]); if (payload.len > 0) { try self.stream.writeAll(payload); } }
pub fn isClosed(self: *const Client) bool { return @atomicLoad(bool, &self._closed, .monotonic); }
fn closeStream(self: *Client) void { if (@atomicRmw(bool, &self._closed, .Xchg, true, .monotonic) == false) { // same lock discipline as close(): never free the stream while a // writer holds (or is queued on) the write lock. self._write_lock.lockUncancelable(self.io); defer self._write_lock.unlock(self.io); self.stream.close(); } }};
// wraps a net.Stream and optional a tls.Clientpub const Stream = struct { io: Io, stream: net.Stream, tls_client: ?*TLSClient = null,
pub fn init(io: Io, stream: net.Stream, tls_client: ?*TLSClient) Stream { return .{ .io = io, .stream = stream, .tls_client = tls_client, }; }
pub fn close(self: *Stream) void { if (self.tls_client) |tls_client| { self.stream.shutdown(self.io, .both) catch {}; tls_client.deinit(); } self.stream.close(self.io); }
pub fn read(self: *Stream, buf: []u8) !usize { if (self.tls_client) |tls_client| { var w: std.Io.Writer = .fixed(buf); while (true) { const n = try tls_client.client.reader.stream(&w, .limited(buf.len)); if (n != 0) { return n; } } } var bufs = [_][]u8{buf}; return self.io.vtable.netRead(self.io.userdata, self.stream.socket.handle, &bufs) catch |err| { return switch (err) { error.ConnectionResetByPeer => error.ConnectionResetByPeer, error.Timeout => error.WouldBlock, else => error.Unexpected, }; }; }
pub fn writeAll(self: *Stream, data: []const u8) !void { if (self.tls_client) |tls_client| { try tls_client.client.writer.writeAll(data); try tls_client.client.writer.flush(); try tls_client.stream_writer.interface.flush(); return; } 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.stream.socket.handle, remaining, &empty, 0) catch |err| { return switch (err) { error.ConnectionResetByPeer => error.ConnectionResetByPeer, else => error.Unexpected, }; }; if (n == 0) return error.Unexpected; remaining = remaining[n..]; } }
const zero_timeout = std.mem.toBytes(posix.timeval{ .sec = 0, .usec = 0 }); pub fn writeTimeout(self: *const Stream, ms: u32) !void { return self.setTimeout(posix.SO.SNDTIMEO, ms); }
pub fn readTimeout(self: *const Stream, ms: u32) !void { return self.setTimeout(posix.SO.RCVTIMEO, ms); }
fn setTimeout(self: *const Stream, opt_name: u32, ms: u32) !void { if (ms == 0) { return self.setsockopt(opt_name, &zero_timeout); }
const timeout = std.mem.toBytes(posix.timeval{ .sec = @intCast(@divTrunc(ms, 1000)), .usec = @intCast(@mod(ms, 1000) * 1000), }); return self.setsockopt(opt_name, &timeout); }
pub fn setsockopt(self: *const Stream, opt_name: u32, value: []const u8) !void { return setConnSockOpt(self.stream.socket.handle, posix.SOL.SOCKET, opt_name, value); }};
/// setsockopt for a *connection* socket.////// std.posix.setsockopt maps BADF/NOTSOCK/INVAL/FAULT to `unreachable`/// ("always a race condition"). On a listening socket that holds. On a/// connection socket it does not: the peer can reset, or another thread can/// close the fd, between connect and the timeout call. `unreachable` cannot be/// caught, so that race aborts the whole process instead of surfacing as an/// error the caller's reconnect path would handle — observed as a panic inside/// the client handshake when an upstream dropped the connection.////// FAULT is deliberately NOT swallowed: a bad option pointer is our own bug,/// not a peer's doing, and hiding it would trade one silent failure for/// another.////// The timeouts this backs are advisory: if the socket really is gone, the/// following read or write reports it properly. So swallow exactly the arms the/// stdlib calls impossible and keep every other error meaningful.////// The server side took the same approach in `setSockOptBestEffort`/// (src/server/server.zig).fn setConnSockOpt(fd: posix.socket_t, level: i32, optname: u32, opt: []const u8) !void { switch (posix.errno(posix.system.setsockopt(fd, level, optname, opt.ptr, @intCast(opt.len)))) { .SUCCESS => {}, // The socket died under us. Not our failure to report. .BADF, .NOTSOCK, .INVAL => {}, .DOM => return error.TimeoutTooBig, .ISCONN => return error.AlreadyConnected, .NOPROTOOPT => return error.InvalidProtocolOption, .NOMEM, .NOBUFS => return error.SystemResources, .PERM => return error.PermissionDenied, .NODEV => return error.NoDevice, .OPNOTSUPP => return error.OperationUnsupported, else => |err| return posix.unexpectedErrno(err), }}
const TLSClient = struct { client: tls.Client, stream: net.Stream, stream_writer: net.Stream.Writer, stream_reader: net.Stream.Reader, arena: std.heap.ArenaAllocator,
fn init(allocator: Allocator, io: Io, stream: net.Stream, config: *const Client.Config) !*TLSClient { var arena = std.heap.ArenaAllocator.init(allocator); errdefer arena.deinit();
const aa = arena.allocator();
const bundle_ptr = try aa.create(Bundle); if (config.ca_bundle) |ca| { bundle_ptr.* = ca; } else { bundle_ptr.* = .{ .map = .empty, .bytes = .empty }; try bundle_ptr.rescan(aa, io, Io.Timestamp.zero); }
const rwlock = try aa.create(std.Io.RwLock); rwlock.* = std.Io.RwLock.init;
// The TLS input and output have to be max_ciphertext_record_len each. // It isn't clear to me how big the un-encrypted reader and writer // need to be. I would think 0, but that will fail an assertion. I // don't think that it's right that we need 4 buffers, but apparently // we do. Until i figure this out, using 4 x max_ciphertext_record_len // seems like the only safe choice. const buf_len = std.crypto.tls.max_ciphertext_record_len; var buf = try aa.alloc(u8, buf_len * 4);
const self = try aa.create(TLSClient); self.* = .{ .stream = stream, .arena = arena, .client = undefined, .stream_writer = stream.writer(io, buf.ptr[0..buf_len][0..buf_len]), .stream_reader = stream.reader(io, buf.ptr[buf_len .. 2 * buf_len][0..buf_len]), };
var entropy: [tls.Client.Options.entropy_len]u8 = undefined; io.random(&entropy); self.client = try tls.Client.init( &self.stream_reader.interface, &self.stream_writer.interface, .{ .ca = .{ .bundle = .{ .gpa = aa, .io = io, .lock = rwlock, .bundle = bundle_ptr } }, .host = .{ .explicit = config.host }, .read_buffer = buf.ptr[2 * buf_len .. 3 * buf_len][0..buf_len], .write_buffer = buf.ptr[3 * buf_len .. 4 * buf_len][0..buf_len], .entropy = &entropy, .realtime_now = Io.Timestamp.now(io, .real), }, );
return self; }
fn deinit(self: *TLSClient) void { _ = self.client.end() catch {}; self.arena.deinit(); }};
fn generateKey(io: Io) [16]u8 { if (comptime @import("builtin").is_test) { return [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; } // 0.16: io.random() fills a buffer var key: [16]u8 = undefined; io.random(&key); return key;}
fn generateMask(io: Io) [4]u8 { // 0.16: io.random() fills a buffer var mask: [4]u8 = undefined; io.random(&mask); return mask;}
fn sendHandshake(path: []const u8, key: []const u8, buf: []u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype) !void { @memcpy(buf[0..4], "GET "); var pos: usize = 4; var end = pos + path.len;
{ @memcpy(buf[pos..end], path); pos = end; }
{ const headers = " HTTP/1.1\r\ncontent-length: 0\r\nupgrade: websocket\r\nsec-websocket-version: 13\r\nconnection: upgrade\r\nsec-websocket-key: "; end = pos + headers.len; @memcpy(buf[pos..end], headers);
pos = end; end = pos + key.len; @memcpy(buf[pos..end], key); }
if (compression) { // NOTE: client_max_window_bits is unsupported const permessage_deflate = "\r\nSec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover; client_no_context_takeover"; pos = end; end = pos + permessage_deflate.len; @memcpy(buf[pos..end], permessage_deflate); }
{ pos = end; end = pos + 2; @memcpy(buf[pos..end], "\r\n"); pos = end; }
if (opts.headers) |extra_headers| { end = pos + extra_headers.len; @memcpy(buf[pos..end], extra_headers); pos = end; if (!std.mem.endsWith(u8, extra_headers, "\r\n")) { buf[pos] = '\r'; buf[pos + 1] = '\n'; pos += 2; } } buf[pos] = '\r'; buf[pos + 1] = '\n';
try stream.writeTimeout(opts.timeout_ms); try stream.writeAll(buf[0 .. pos + 2]); try stream.writeTimeout(0);}
const HandShakeReply = struct { compression: bool, over_read: usize,
/// On a non-101 status line, read the rest of the rejection (bounded by /// buf and the already-armed read timeout) and capture the status plus /// up to body_buf.len bytes of body. Returns null when the status line /// is too short/unparseable — that is a malformed response, not a /// rejection. Content-Length, when present and sane, bounds the body /// read exactly; otherwise reading stops at close, timeout, or the cap. fn captureRejection(status_line: []const u8, buf: []u8, initial_pos: usize, stream: anytype) ?Client.HandshakeFailure { if (status_line.len < 12 or status_line[8] != ' ') return null; const status = std.fmt.parseInt(u16, status_line[9..12], 10) catch return null; var f: Client.HandshakeFailure = .{ .status = status };
var pos = initial_pos; const body_start = while (true) { if (std.mem.indexOf(u8, buf[0..pos], "\r\n\r\n")) |i| break i + 4; if (pos == buf.len) return f; const n = stream.read(buf[pos..]) catch return f; if (n == 0) return f; pos += n; };
// Only read further when Content-Length promises the bytes exist: a // speculative read on an open connection blocks until the handshake // timeout (and SO_RCVTIMEO's EAGAIN panics under Io.Threaded in // debug). Without Content-Length, the body is whatever arrived with // the headers. if (findContentLength(buf[0..body_start])) |cl| { const want = @min(cl, f.body_buf.len); while (pos - body_start < want and pos < buf.len) { const n = stream.read(buf[pos..]) catch break; if (n == 0) break; pos += n; } } const len = @min(pos - body_start, f.body_buf.len); @memcpy(f.body_buf[0..len], buf[body_start .. body_start + len]); f.body_len = len; return f; }
fn findContentLength(headers: []const u8) ?usize { var it = std.mem.splitSequence(u8, headers, "\r\n"); while (it.next()) |line| { const colon = std.mem.indexOfScalar(u8, line, ':') orelse continue; if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[0..colon], " "), "content-length")) continue; return std.fmt.parseInt(usize, std.mem.trim(u8, line[colon + 1 ..], " "), 10) catch null; } return null; }
fn read(buf: []u8, key: []const u8, opts: *const Client.HandshakeOpts, compression: bool, stream: anytype, failure: *?Client.HandshakeFailure) !HandShakeReply { const timeout_ms = opts.timeout_ms; const deadline = milliTimestamp(stream.io) + timeout_ms; try stream.readTimeout(timeout_ms);
var pos: usize = 0; var line_start: usize = 0; var complete_response: u8 = 0; var server_compression: bool = false;
while (true) { // 0.16: using libc recv, WouldBlock indicates timeout const n = stream.read(buf[pos..]) catch |err| switch (err) { error.WouldBlock => return error.Timeout, else => return err, }; if (n == 0) { return error.ConnectionClosed; }
pos += n; while (std.mem.indexOfScalar(u8, buf[line_start..pos], '\r')) |relative_end| { if (relative_end == 0) { if (complete_response != 15) { return error.InvalidHandshakeResponse; } // TCP can split the terminating CRLF — if the trailing \n // hasn't arrived yet, pos == line_start + 1 and the over_read // subtraction below would underflow. break for more data. if (line_start + 2 > pos) break; const over_read = pos - (line_start + 2); std.mem.copyForwards(u8, buf[0..over_read], buf[line_start + 2 .. pos]); try stream.readTimeout(0); return .{ .over_read = over_read, .compression = server_compression, }; }
const line_end = line_start + relative_end; const line = buf[line_start..line_end];
// the next line starts where this line ends, skip over the \r\n. // TCP can split mid-CRLF — if \n hasn't arrived yet, break to // the outer read loop for more data. line_start = line_end + 2; if (line_start > pos) break;
if (complete_response == 0) { if (!ascii.startsWithIgnoreCase(line, "HTTP/1.1 101 ")) { // A well-formed non-101 response is a server-side // rejection with meaning: capture the status and a // bounded body so the caller can act on it (see // Client.handshake_failure), then fail as before. failure.* = captureRejection(line, buf, pos, stream); return error.InvalidHandshakeResponse; } complete_response |= 1; continue; }
for (line, 0..) |b, i| { // find the colon and lowercase the header while we're iterating if ('A' <= b and b <= 'Z') { line[i] = b + 32; continue; }
if (b != ':') { continue; }
switch (i) { 7 => if (std.mem.eql(u8, line[0..i], "upgrade")) { if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "websocket")) { return error.InvalidUpgradeHeader; } complete_response |= 2; }, 10 => if (std.mem.eql(u8, line[0..i], "connection")) { if (!ascii.eqlIgnoreCase(std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace), "upgrade")) { return error.InvalidConnectionHeader; } complete_response |= 4; }, 20 => if (std.mem.eql(u8, line[0..i], "sec-websocket-accept")) { var h: [20]u8 = undefined; { var hasher = std.crypto.hash.Sha1.init(.{}); hasher.update(key); hasher.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); hasher.final(&h); }
var encoded_buf: [28]u8 = undefined; const sec_hash = std.base64.standard.Encoder.encode(&encoded_buf, &h); const header_value = std.mem.trim(u8, line[i + 1 ..], &ascii.whitespace);
if (!std.mem.eql(u8, header_value, sec_hash)) { return error.InvalidWebsocketAcceptHeader; } complete_response |= 8; }, 24 => if (std.mem.eql(u8, line[0..i], "sec-websocket-extensions")) { if (try parseExtension(line[i + 1 ..])) |sc| { if (!compression) { // server is saying compression, but we didn't ask for it. return error.InvalidExtensionHeader; } if (!sc.client_no_context_takeover or !sc.server_no_context_takeover) { // as of Zig 0.15, we no longer support context takeover // We told the server this, it should have respected it. return error.InvalidExtensionHeader; }
server_compression = true; } }, else => {}, // some other header we don't care about } } }
if (milliTimestamp(stream.io) > deadline) { return error.Timeout; }
if (pos == buf.len) { return error.ResponseTooLarge; } } }
pub fn parseExtension(value: []const u8) !?ServerHandshake.Compression { var deflate = false; var client_max_bits: u8 = 15; var client_no_context_takeover = false; var server_no_context_takeover = false;
var it = std.mem.splitScalar(u8, value, ';'); while (it.next()) |param_| { const param = std.mem.trim(u8, param_, &ascii.whitespace); if (std.mem.eql(u8, param, "permessage-deflate")) { deflate = true; continue; } if (std.mem.eql(u8, param, "client_no_context_takeover")) { client_no_context_takeover = true; continue; } if (std.mem.eql(u8, param, "server_no_context_takeover")) { server_no_context_takeover = true; continue; } const client_max_window_bits = "client_max_window_bits="; if (std.mem.startsWith(u8, param, client_max_window_bits)) { client_max_bits = std.fmt.parseInt(u8, param[client_max_window_bits.len..], 10) catch { return error.InvalidCompressionServerMaxBits; }; } } if (deflate == false) { return null; }
if (client_max_bits != 15) { // We don't offer client window, so if the server asks for one, that's an error return error.InvalidExtensionHeader; }
return .{ .client_no_context_takeover = client_no_context_takeover, .server_no_context_takeover = server_no_context_takeover, }; }};
const t = @import("../t.zig");test "Client: handshake" { { // empty response var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); }
{ // invalid websocket response var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("HTTP/1.1 200 OK\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); // a well-formed non-101 is a rejection: status captured, empty body try t.expectEqual(@as(u16, 200), client.handshake_failure.?.status); try t.expectEqual(@as(usize, 0), client.handshake_failure.?.body().len); }
{ // pre-upgrade rejection with a JSON error envelope (the jetstream // subscribeEvents contract: 400 CursorTooOld etc.) — the caller // needs the status and body to pick terminal vs transient handling var pair = t.SocketPair.init(.{}); defer pair.deinit(); const body = "{\"error\":\"CursorTooOld\",\"message\":\"below floor\"}"; var response_buf: [256]u8 = undefined; const response = std.fmt.bufPrint( &response_buf, "HTTP/1.1 400 Bad Request\r\nContent-Type: application/json\r\nContent-Length: {d}\r\n\r\n{s}", .{ body.len, body }, ) catch unreachable; try pair.clientWriteAll(response);
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); const failure = client.handshake_failure.?; try t.expectEqual(@as(u16, 400), failure.status); try t.expectString(body, failure.body()); }
{ // missing upgrade header var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); }
{ // wrong upgrade header var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: nope\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidUpgradeHeader, client.handshake("/", .{})); }
{ // missing connection header var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); }
{ // wrong connection header var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: something\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidConnectionHeader, client.handshake("/", .{})); }
{ // missing Sec-Websocket-Accept header var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nUpgrade: websocket\r\nConnection: upgrade\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidHandshakeResponse, client.handshake("/", .{})); }
{ // wrong Sec-Websocket-Accept header var pair = t.SocketPair.init(.{}); defer pair.deinit(); try pair.clientWriteAll("HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: hack\r\n\r\n");
var client = testClient(&pair); defer client.deinit(); try t.expectError(error.InvalidWebsocketAcceptHeader, client.handshake("/", .{})); }
{ // ok for successful var pair = t.SocketPair.init(.{}); defer pair.deinit(); 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");
var client = testClient(&pair); defer client.deinit(); try client.handshake("/", .{}); try t.expectEqual(0, client._reader.pos); }
{ // ok for successful, with overread var pair = t.SocketPair.init(.{}); defer pair.deinit(); 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");
var client = testClient(&pair); defer client.deinit(); try client.handshake("/", .{}); try t.expectEqual(50, client._reader.pos); }}
test "Client: handshake with terminating CRLF split across reads" { // regression: when TCP delivers the final \r of the blank-line CRLF as the // last byte of a read and the \n arrives in a later read, the end-of-headers // branch computed `over_read = pos - (line_start + 2)` while pos == line_start // + 1, underflowing usize → integer-overflow panic. parser must wait for \n. const io = std.Options.debug_io;
// generateKey() is deterministic under test ({1..16}); this accept matches it. const bin_key = [16]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16 }; var encoded_key: [24]u8 = undefined; const key = std.base64.standard.Encoder.encode(&encoded_key, &bin_key);
const headers = "HTTP/1.1 101 Switching Protocol\r\nupgrade: WebSocket\r\nConnection: UPGRADE\r\nSec-Websocket-Accept: C/0nmHhBztSRGR1CwL6Tf4ZjwpY=\r\n"; const trailing = "Some Random Data Which is Part Of the Next Message";
const ChunkStream = struct { io: Io, chunks: []const []const u8, idx: usize = 0, fn read(self: *@This(), buf: []u8) !usize { if (self.idx >= self.chunks.len) return 0; const c = self.chunks[self.idx]; std.debug.assert(c.len <= buf.len); @memcpy(buf[0..c.len], c); self.idx += 1; return c.len; } fn readTimeout(self: *const @This(), ms: u32) !void { _ = self; _ = ms; } };
// first read ends on the terminating \r; the \n (and over-read) arrive next. var chunks = [_][]const u8{ headers ++ "\r", "\n" ++ trailing }; var stream = ChunkStream{ .io = io, .chunks = &chunks };
var buf: [4096]u8 = undefined; const opts = Client.HandshakeOpts{}; var failure: ?Client.HandshakeFailure = null; const res = try HandShakeReply.read(&buf, key, &opts, false, &stream, &failure); try t.expectEqual(trailing.len, res.over_read); try t.expectSlice(u8, trailing, buf[0..res.over_read]);}
test "Client: setting a timeout on a dead socket does not abort the process" { // regression: std.posix.setsockopt maps BADF/NOTSOCK/INVAL/FAULT to // `unreachable`. Those are reachable on a *connection* socket — the peer // resets, or another thread closes the fd, between connect and the timeout // call. `unreachable` cannot be caught, so an upstream dropping the // connection aborted the whole process from inside the handshake instead of // returning an error the caller could reconnect on. const io = std.Options.debug_io; const stream = try testConnectedStream(io, "127.0.0.1", 9292); var dead = Stream.init(io, stream, null); // Close underneath the Stream: every later setsockopt sees EBADF, which is // precisely the race the stdlib declares impossible. stream.close(io);
// Before the fix each of these panicked rather than returning. try dead.readTimeout(5_000); try dead.writeTimeout(5_000); try dead.readTimeout(0);}
test "Client: writeFrame after close returns error.Closed, never touches the stream" { // regression: close() released the write lock between sending the close // frame and stream.close(), so a concurrent writer queued on the lock // (zlay's keepalive pinger) resumed into a freed TLS client — a GPF in // production on 2026-08-27. writers must observe _closed under the lock. var pair = t.SocketPair.init(.{}); defer pair.deinit();
var client = testClient(&pair); defer client.deinit();
var payload = [_]u8{ 'h', 'i' }; try client.writeFrame(.text, &payload);
try client.close(.{ .code = 4000 }); try t.expectError(error.Closed, client.writePing(&payload)); try t.expectError(error.Closed, client.writeFrame(.text, &payload)); // idempotent close still fine after the lock-scoped teardown try client.close(.{});}
test "Client: write/read" { const io = std.Options.debug_io; const stream = try testConnectedStream(io, "127.0.0.1", 9292); var client = try Client.initWithStream(io, t.allocator, stream, .{ .port = 9292, .host = "127.0.0.1", }); defer client.deinit();
try client.handshake("/", .{ .timeout_ms = 1000, });
var buf = [_]u8{ 'o', 'v', 'e', 'r' }; try client.write(&buf); try client.readTimeout(1000);
const message = (try client.read()) orelse unreachable; try t.expectEqual(.text, message.type); try t.expectString("9000", message.data);
client.close(.{}) catch unreachable;}
test "Client: close with code" { const io = std.Options.debug_io; const stream = try testConnectedStream(io, "127.0.0.1", 9292); var client = try Client.initWithStream(io, t.allocator, stream, .{ .port = 9292, .host = "127.0.0.1", }); defer client.deinit();
try client.handshake("/", .{ .timeout_ms = 1000, });
client.close(.{ .code = 4002 }) catch unreachable;}
test "Client: with code and reason" { const io = std.Options.debug_io; const stream = try testConnectedStream(io, "127.0.0.1", 9292); var client = try Client.initWithStream(io, t.allocator, stream, .{ .port = 9292, .host = "127.0.0.1", }); defer client.deinit();
try client.handshake("/", .{ .timeout_ms = 1000, });
client.close(.{ .code = 4002, .reason = "goodbye" }) catch unreachable;}
test "Client: Handler" { var h = try ClientHandler.init(t.allocator); defer h.deinit();
var buf: [6]u8 = undefined; { @memcpy(buf[0..3], "dyn"); try h.client.write(buf[0..3]); }
{ @memcpy(buf[0..4], "ping"); try h.client.write(buf[0..4]); }
{ @memcpy(buf[0..4], "pong"); try h.client.write(buf[0..4]); }
{ @memcpy(buf[0..6], "close1"); try h.client.write(buf[0..6]); }
try h.client.readLoop(&h);
// if pong is true then ping and message have to be true // because each asserts the previous try t.expectEqual(true, h.pong); try t.expectEqual(true, h.closed);}
fn testClient(pair: *t.SocketPair) Client { pair.server_taken = true; const stream = pair.server; const io = std.Options.debug_io; const bp = t.allocator.create(buffer.Provider) catch unreachable; bp.* = buffer.Provider.init(t.allocator, .{ .count = 0, .size = 0, .max = 4096 }) catch unreachable;
const reader_buf = bp.allocator.alloc(u8, 1024) catch unreachable;
return .{ .io = io, ._closed = false, ._own_bp = true, ._mask_fn = generateMask, ._compression_opts = null, .stream = .{ .io = io, .stream = stream }, ._reader = Reader.init(reader_buf, bp, null), };}
fn testConnectedStream(io: Io, host: []const u8, port: u16) !net.Stream { const addr = try net.IpAddress.parse(host, port); return net.IpAddress.connect(&addr, io, .{ .mode = .stream });}
const ClientHandler = struct { ping: bool = false, pong: bool = false, closed: bool = false, message: bool = false, client: Client,
fn init(allocator: Allocator) !ClientHandler { const io = std.Options.debug_io; const stream = try testConnectedStream(io, "127.0.0.1", 9292); var client = try Client.initWithStream(io, allocator, stream, .{ .port = 9292, .host = "127.0.0.1", }); errdefer client.deinit();
try client.handshake("/", .{ .timeout_ms = 1000, });
return .{ .client = client, }; }
fn deinit(self: *ClientHandler) void { self.client.deinit(); }
pub fn serverMessage(self: *ClientHandler, data: []u8, tpe: proto.Message.TextType) !void { try t.expectEqual(.text, tpe); try t.expectString("over 9000!", data); self.message = true; }
pub fn serverPing(self: *ClientHandler, data: []u8) !void { try t.expectEqual(true, self.message); try t.expectString("a-ping", data); self.ping = true; }
pub fn serverPong(self: *ClientHandler, data: []u8) !void { try t.expectEqual(true, self.ping); try t.expectString("a-pong", data); self.pong = true; }
pub fn close(self: *ClientHandler) void { self.client.close(.{}) catch unreachable; self.closed = true; }};