tiny.http.WebSocket
Defined in tiny.http.
API (24)
Actions
Public operations.
activatecapacityclosedeinithandshakeinitreceivesendBinarysendPingsendPongsendTextvalidateClientHandshakevalidateClientKey
Types and contracts
Public types and contracts.
Fields and members
Public fields and members.
connectionfragment_lengthfragment_opcodefragment_returnedpending_consumedread_lengthstatestorage
Source
Source: lib/http/src/websocket.zig:564
zig
pub const WebSocket = struct { connection: *Connection, state: WebSocketState, storage: SessionStorage, read_length: usize = 0, pending_consumed: usize = 0, fragment_opcode: ?Opcode, fragment_length: usize = 0, fragment_returned: bool = false, pub const Limits: type = SessionLimits; pub const Capacity: type = SessionCapacity; pub const Storage: type = SessionStorage; pub fn init( allocator: std.mem.Allocator, limits: SessionLimits, connection: *Connection, ) SessionStorage.InitError!WebSocket { return .{ .connection = connection, .state = .open, .storage = try SessionStorage.init(allocator, limits), .fragment_opcode = null, }; } pub fn activate(self: *WebSocket) void { self.storage.activate(); } pub fn deinit(self: *WebSocket, allocator: std.mem.Allocator) void { self.storage.deinit(allocator); self.* = undefined; } pub fn capacity(self: *const WebSocket) SessionCapacity { return self.storage.capacity; } pub fn validateClientHandshake(req: *const Request) HandshakeError!WebSocketClientHandshake { if (req.method != .GET) return HandshakeError.InvalidMethod; if (req.version != .http_1_1) return HandshakeError.InvalidHttpVersion; if (!req.isWebSocketUpgrade()) return HandshakeError.InvalidUpgrade; const version = req.headers.get("Sec-WebSocket-Version") orelse return HandshakeError.MissingVersion; if (!std.mem.eql(u8, version, "13")) return HandshakeError.UnsupportedVersion; const key = req.getWebSocketKey() orelse return HandshakeError.MissingKey; try validateClientKey(key); return .{ .key = key }; } pub fn validateClientKey(key: []const u8) HandshakeError!void { const decoded_len = std.base64.standard.Decoder.calcSizeForSlice(key) catch return HandshakeError.InvalidKey; if (decoded_len != 16) return HandshakeError.InvalidKey; var decoded: [16]u8 = undefined; std.base64.standard.Decoder.decode(&decoded, key) catch return HandshakeError.InvalidKey; } pub fn handshake(key: []const u8) HandshakeError![28]u8 { try validateClientKey(key); return accept_key(key); } pub fn sendText(self: *WebSocket, data: []const u8) !void { try self.sendFrame(.{ .fin = true, .opcode = .text, .mask = null, .payload = data, }); } pub fn sendBinary(self: *WebSocket, data: []const u8) !void { try self.sendFrame(.{ .fin = true, .opcode = .binary, .mask = null, .payload = data, }); } pub fn sendPing(self: *WebSocket, data: []const u8) !void { try self.sendFrame(.{ .fin = true, .opcode = .ping, .mask = null, .payload = data, }); } pub fn sendPong(self: *WebSocket, data: []const u8) !void { try self.sendFrame(.{ .fin = true, .opcode = .pong, .mask = null, .payload = data, }); } pub fn close(self: *WebSocket, code: u16, reason: []const u8) !void { if (self.state == .closed) return; const payload_len = std.math.add(usize, 2, reason.len) catch { return FrameError.InvalidFrame; }; if (payload_len > maximum_control_payload_bytes) return FrameError.InvalidFrame; var payload: [maximum_control_payload_bytes]u8 = undefined; std.mem.writeInt(u16, payload[0..2], code, .big); @memcpy(payload[2..][0..reason.len], reason); try self.sendFrame(.{ .fin = true, .opcode = .close, .mask = null, .payload = payload[0..payload_len], }); self.state = .closing; } pub fn receive(self: *WebSocket) !?WebSocketMessage { std.debug.assert(self.storage.phase == .steady); self.releaseBorrowedMessage(); while (self.state != .closed) { const buffered_length = self.read_length; if (try self.processBufferedFrame()) |message| return message; if (self.state == .closed) return null; if (self.read_length != buffered_length) continue; try self.readFrameBytes(); } return null; } fn processBufferedFrame(self: *WebSocket) !?WebSocketMessage { if (self.read_length == 0) return null; const frame_region = self.frameRegion(); const parsed = Frame.parseServer( frame_region[0..self.read_length], self.storage.capacity.frame_payload_bytes, ) catch |err| switch (err) { error.IncompleteFrame => return null, error.MaskRequired, error.InvalidOpcode, error.InvalidFrame => { self.rejectProtocol(); return null; }, error.PayloadTooLarge => { self.rejectMessageTooBig(); return null; }, else => return err, }; if (parsed.frame.opcode.isControl()) { try self.processControlFrame(parsed.frame, parsed.consumed); return null; } return self.processDataFrame(parsed.frame, parsed.consumed); } fn processControlFrame(self: *WebSocket, frame: Frame, consumed: usize) !void { defer self.consumeFrame(consumed); switch (frame.opcode) { .ping => try self.sendPong(frame.payload), .pong => {}, .close => { if (self.state != .closing) self.replyToClose(frame.payload); self.state = .closed; self.connection.markClosing(); }, else => unreachable, } } fn replyToClose(self: *WebSocket, payload: []const u8) void { if (payload.len >= 2) { const code = std.mem.readInt(u16, payload[0..2], .big); self.close(code, payload[2..]) catch |err| { log.warn("failed to send close response: {s}", .{@errorName(err)}); }; } else { self.close(WebSocketCloseCode.normal, "") catch |err| { log.warn("failed to send close response: {s}", .{@errorName(err)}); }; } } fn processDataFrame(self: *WebSocket, frame: Frame, consumed: usize) ?WebSocketMessage { if (frame.fin) return self.finishDataFrame(frame, consumed); if (frame.opcode == .continuation) { if (self.fragment_opcode == null) { self.rejectProtocol(); return null; } } else if (self.fragment_opcode == null) { self.fragment_opcode = frame.opcode; self.fragment_length = 0; } else { self.rejectProtocol(); return null; } if (!self.retainFragment(frame.payload)) return null; self.consumeFrame(consumed); return null; } fn finishDataFrame(self: *WebSocket, frame: Frame, consumed: usize) ?WebSocketMessage { if (self.fragment_opcode) |opcode| { if (frame.opcode != .continuation) { self.rejectProtocol(); return null; } if (!self.retainFragment(frame.payload)) return null; self.consumeFrame(consumed); self.fragment_opcode = null; self.fragment_returned = true; return .{ .opcode = opcode, .payload = self.messageRegion()[0..self.fragment_length], }; } if (frame.opcode == .continuation) { self.rejectProtocol(); return null; } if (frame.payload.len > self.storage.capacity.message_payload_bytes) { self.rejectMessageTooBig(); return null; } self.pending_consumed = consumed; return .{ .opcode = frame.opcode, .payload = frame.payload }; } fn retainFragment(self: *WebSocket, payload: []const u8) bool { if (!fragmentLengthAllowed( self.fragment_length, payload.len, self.storage.capacity.message_payload_bytes, )) { self.rejectMessageTooBig(); return false; } self.appendFragment(payload); return true; } fn readFrameBytes(self: *WebSocket) !void { const frame_region = self.frameRegion(); if (self.read_length == frame_region.len) { self.rejectMessageTooBig(); return; } const read = self.connection.read( frame_region[self.read_length..], ) catch |err| switch (err) { error.WouldBlock => return error.WouldBlock, error.ConnectionClosed => { self.state = .closed; self.connection.markClosing(); return; }, else => return err, }; if (read == 0) { self.state = .closed; self.connection.markClosing(); return; } self.read_length += read; } fn sendFrame(self: *WebSocket, frame: Frame) !void { std.debug.assert(self.storage.phase == .steady); std.debug.assert(frame.mask == null); var header: [maximum_frame_header_bytes]u8 = undefined; const encoded = try frame.writeHeader(&header, std.math.maxInt(usize)); try self.connection.write(encoded); if (frame.payload.len != 0) try self.connection.write(frame.payload); } fn frameRegion(self: *WebSocket) []u8 { return self.storage.frame(self.storage.capacity.frame_payload_bytes) catch unreachable; } fn messageRegion(self: *WebSocket) []u8 { return self.storage.message(self.storage.capacity.message_payload_bytes) catch unreachable; } fn appendFragment(self: *WebSocket, payload: []const u8) void { const message = self.messageRegion(); @memcpy(message[self.fragment_length..][0..payload.len], payload); self.fragment_length += payload.len; } fn consumeFrame(self: *WebSocket, consumed: usize) void { std.debug.assert(consumed <= self.read_length); const remaining = self.read_length - consumed; const frame = self.frameRegion(); std.mem.copyForwards(u8, frame[0..remaining], frame[consumed..self.read_length]); self.read_length = remaining; } fn releaseBorrowedMessage(self: *WebSocket) void { if (self.pending_consumed != 0) { self.consumeFrame(self.pending_consumed); self.pending_consumed = 0; } if (self.fragment_returned) { self.fragment_length = 0; self.fragment_returned = false; } } fn rejectProtocol(self: *WebSocket) void { self.close(WebSocketCloseCode.protocol_error, "Invalid WebSocket frame") catch |err| { log.warn("failed to send protocol error close: {s}", .{@errorName(err)}); }; self.state = .closed; self.connection.markClosing(); } fn rejectMessageTooBig(self: *WebSocket) void { self.close(WebSocketCloseCode.message_too_big, "Message too big") catch |close_err| { log.warn("failed to send message-too-big close: {s}", .{@errorName(close_err)}); }; self.state = .closed; self.connection.markClosing(); }};Source: lib/http/src/root.zig:61
zig
pub const WebSocket = websocket.WebSocket;Audit
| Definitions | 17 |
|---|---|
| Public names | 17 |
| Members | 8 |
| Version | 26.7.0 |
| Revision | daab053ee433 |