diff options
Diffstat (limited to 'src/http.zig')
| -rw-r--r-- | src/http.zig | 339 |
1 files changed, 339 insertions, 0 deletions
diff --git a/src/http.zig b/src/http.zig new file mode 100644 index 0000000..0a07160 --- /dev/null +++ b/src/http.zig @@ -0,0 +1,339 @@ +//! 9P over HTTP WebSockets. The framing layer accepts arbitrary std.Io readers +//! and writers; it does not own sockets, TLS, UARTs, filesystems, or mounts. +//! HTTP/1.1 connection setup uses std.http.Client (including verified HTTPS). +const std = @import("std"); +const wire = @import("wire.zig"); +const Io = std.Io; +const guid = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; + +pub const Opcode = enum(u4) { continuation = 0, binary = 2, close = 8, ping = 9, pong = 10, _ }; +pub const Role = enum { client, server }; +pub const Message = struct { opcode: Opcode, data: []const u8 }; + +/// Caller owns both streams and serializes writes. One read and one write may +/// run concurrently. Keep the receive buffer stable across interleaved controls. +/// On any protocol/I/O error, discard the WebSocket and close its underlying stream. +pub const WebSocket = struct { + input: *Io.Reader, + output: *Io.Writer, + role: Role, + fragmented: bool = false, + used: usize = 0, + control: [125]u8 = undefined, + + /// One complete binary message contains exactly one 9P frame. Fragmented + /// WebSocket messages and interleaved control frames are supported. + pub fn receive(s: *WebSocket, buffer: []u8) !Message { + while (true) { + var header: [2]u8 = undefined; + try s.input.readSliceAll(&header); + const fin = header[0] & 0x80 != 0; + if (header[0] & 0x70 != 0) return error.ReservedBits; + const opcode: Opcode = @enumFromInt(header[0] & 0x0f); + const control = header[0] & 8 != 0; + const masked = header[1] & 0x80 != 0; + if (masked != (s.role == .server)) return error.InvalidMask; + const short = header[1] & 0x7f; + const size: u64 = switch (short) { + 126 => value: { + const n = try readInt(s.input, u16); + if (n < 126) return error.NonCanonicalLength; + break :value n; + }, + 127 => value: { + const n = try readInt(s.input, u64); + if (n < 65536 or n >> 63 != 0) return error.NonCanonicalLength; + break :value n; + }, + else => short, + }; + if (control and (!fin or size > 125)) return error.InvalidControl; + switch (opcode) { + .binary => if (s.fragmented) return error.ExpectedContinuation, + .continuation => if (!s.fragmented) return error.UnexpectedContinuation, + .close, .ping, .pong => {}, + else => return error.ExpectedBinary, + } + var mask: [4]u8 = @splat(0); + if (masked) try s.input.readSliceAll(&mask); + const dest = if (control) s.control[0..] else buffer[s.used..]; + if (size > dest.len) return error.MessageTooLarge; + const payload = dest[0..@intCast(size)]; + try s.input.readSliceAll(payload); + if (masked) for (payload, 0..) |*byte, i| { + byte.* ^= mask[i % 4]; + }; + if (control) { + if (opcode == .close) try validateClose(payload); + return .{ .opcode = opcode, .data = payload }; + } + s.used += payload.len; + s.fragmented = !fin; + if (fin) { + const frame = buffer[0..s.used]; + s.used = 0; + _ = try wire.decode(frame); + return .{ .opcode = .binary, .data = frame }; + } + } + } + + /// Clients must supply a fresh, unpredictable mask for every message. + /// Servers pass null. This call flushes the supplied writer. + pub fn send(s: *WebSocket, bytes: []const u8, opcode: Opcode, mask: ?[4]u8) !void { + try s.sendUnflushed(bytes, opcode, mask); + try s.output.flush(); + } + + /// For layered writers whose outer connection controls flushing (e.g. TLS). + pub fn sendUnflushed(s: *WebSocket, bytes: []const u8, opcode: Opcode, mask: ?[4]u8) !void { + if ((mask != null) != (s.role == .client)) return error.InvalidMask; + switch (opcode) { + .binary => { + _ = try wire.decode(bytes); + }, + .close => { + if (bytes.len > 125) return error.InvalidControl; + try validateClose(bytes); + }, + .ping, .pong => if (bytes.len > 125) return error.InvalidControl, + else => return error.ExpectedBinary, + } + const out = s.output; + try out.writeByte(0x80 | @as(u8, @intFromEnum(opcode))); + const bit: u8 = if (mask != null) 0x80 else 0; + if (bytes.len < 126) { + try out.writeByte(bit | @as(u8, @intCast(bytes.len))); + } else if (bytes.len <= 65535) { + try out.writeByte(bit | 126); + try out.writeInt(u16, @intCast(bytes.len), .big); + } else { + try out.writeByte(bit | 127); + try out.writeInt(u64, bytes.len, .big); + } + if (mask) |key| { + try out.writeAll(&key); + var scratch: [1024]u8 = undefined; + var offset: usize = 0; + while (offset < bytes.len) { + const count = @min(scratch.len, bytes.len - offset); + for (scratch[0..count], bytes[offset..][0..count], 0..) |*to, from, i| to.* = from ^ key[(offset + i) % 4]; + try out.writeAll(scratch[0..count]); + offset += count; + } + } else try out.writeAll(bytes); + } +}; + +fn readInt(reader: *Io.Reader, comptime T: type) !T { + var bytes: [@sizeOf(T)]u8 = undefined; + try reader.readSliceAll(&bytes); + return std.mem.readInt(T, &bytes, .big); +} + +fn validateClose(bytes: []const u8) !void { + if (bytes.len == 1) return error.InvalidClose; + if (bytes.len < 2) return; + const code = std.mem.readInt(u16, bytes[0..2], .big); + if (code < 1000 or code >= 5000 or code == 1004 or code == 1005 or code == 1006 or (code >= 1015 and code < 3000)) return error.InvalidClose; + if (!std.unicode.utf8ValidateSlice(bytes[2..])) return error.InvalidClose; +} + +fn hasToken(value: []const u8, token: []const u8) bool { + var it = std.mem.splitScalar(u8, value, ','); + while (it.next()) |part| if (std.ascii.eqlIgnoreCase(std.mem.trim(u8, part, " \t"), token)) return true; + return false; +} + +/// Validate and accept an HTTP/1.1 upgrade. The application must check the +/// request route, Host, Origin, and authorization before calling this function. +pub fn accept(request: *std.http.Server.Request) !WebSocket { + var connection = false; + var version = false; + var it = request.iterateHeaders(); + while (it.next()) |header| { + if (std.ascii.eqlIgnoreCase(header.name, "connection")) connection = hasToken(header.value, "upgrade"); + if (std.ascii.eqlIgnoreCase(header.name, "sec-websocket-version")) version = std.mem.eql(u8, header.value, "13"); + } + if (!connection or !version or (request.head.content_length orelse 0) != 0 or request.head.transfer_encoding != .none) return error.InvalidUpgrade; + const key = switch (request.upgradeRequested()) { + .websocket => |value| value orelse return error.InvalidUpgrade, + else => return error.InvalidUpgrade, + }; + if (key.len != 24) return error.InvalidUpgrade; + const nonce_len = std.base64.standard.Decoder.calcSizeForSlice(key) catch return error.InvalidUpgrade; + if (nonce_len != 16) return error.InvalidUpgrade; + var nonce: [16]u8 = undefined; + std.base64.standard.Decoder.decode(&nonce, key) catch return error.InvalidUpgrade; + var ws = try request.respondWebSocket(.{ .key = key }); + try ws.flush(); + return .{ .input = ws.input, .output = ws.output, .role = .server }; +} + +/// Owns one upgraded connection in the caller's std.http.Client. Both that +/// client and the URL bytes must outlive this value. deinit closes the connection; +/// upgraded connections are never returned to the HTTP keep-alive pool. +pub const Client = struct { + request: std.http.Client.Request, + socket: WebSocket, + + pub fn connect(http: *std.http.Client, url: []const u8, origin: []const u8) !Client { + const uri = try std.Uri.parse(url); + if (!std.mem.eql(u8, uri.scheme, "http") and !std.mem.eql(u8, uri.scheme, "https")) return error.UnsupportedUriScheme; + if (std.mem.findScalar(u8, origin, '\r') != null or std.mem.findScalar(u8, origin, '\n') != null) return error.InvalidOrigin; + var nonce: [16]u8 = undefined; + try http.io.randomSecure(&nonce); + var key: [24]u8 = undefined; + _ = std.base64.standard.Encoder.encode(&key, &nonce); + const headers: []const std.http.Header = &.{ + .{ .name = "upgrade", .value = "websocket" }, + .{ .name = "sec-websocket-version", .value = "13" }, + .{ .name = "sec-websocket-key", .value = &key }, + .{ .name = "origin", .value = origin }, + }; + var request = try http.request(.GET, uri, .{ + .redirect_behavior = .unhandled, + .headers = .{ .connection = .{ .override = "Upgrade" }, .accept_encoding = .omit }, + .extra_headers = headers, + }); + errdefer { + request.connection.?.closing = true; + request.deinit(); + } + try request.sendBodiless(); + const response = try request.receiveHead(&.{}); + request.connection.?.closing = true; + if (response.head.status != .switching_protocols) return error.UpgradeRejected; + var sha = std.crypto.hash.Sha1.init(.{}); + sha.update(&key); + sha.update(guid); + var digest: [20]u8 = undefined; + sha.final(&digest); + var expected: [28]u8 = undefined; + _ = std.base64.standard.Encoder.encode(&expected, &digest); + var accepted = false; + var upgraded = false; + var connection = false; + var it = response.head.iterateHeaders(); + while (it.next()) |header| { + if (std.ascii.eqlIgnoreCase(header.name, "sec-websocket-accept")) accepted = std.mem.eql(u8, header.value, &expected); + if (std.ascii.eqlIgnoreCase(header.name, "upgrade")) upgraded = std.ascii.eqlIgnoreCase(header.value, "websocket"); + if (std.ascii.eqlIgnoreCase(header.name, "connection")) connection = hasToken(header.value, "upgrade"); + if (std.ascii.eqlIgnoreCase(header.name, "sec-websocket-extensions") or std.ascii.eqlIgnoreCase(header.name, "sec-websocket-protocol")) return error.UnsupportedExtension; + } + if (!accepted or !upgraded or !connection) return error.InvalidUpgrade; + request.extra_headers = &.{}; + const stream = request.connection.?; + return .{ .request = request, .socket = .{ .input = stream.reader(), .output = stream.writer(), .role = .client } }; + } + + pub fn deinit(client: *Client) void { + client.request.deinit(); + client.* = undefined; + } + + pub fn send(client: *Client, frame: []const u8) !void { + try client.sendMessage(frame, .binary); + } + + fn sendMessage(client: *Client, bytes: []const u8, opcode: Opcode) !void { + var mask: [4]u8 = undefined; + try client.request.client.io.randomSecure(&mask); + try client.socket.sendUnflushed(bytes, opcode, mask); + try client.request.connection.?.flush(); + } + + /// Serial convenience API. Handles ping/pong and close; use socket.receive + /// and separately serialized writes when driving a concurrent session. + pub fn receive(client: *Client, buffer: []u8) ![]const u8 { + while (true) { + const message = try client.socket.receive(buffer); + switch (message.opcode) { + .binary => return message.data, + .ping => try client.sendMessage(message.data, .pong), + .pong => {}, + .close => { + try client.sendMessage(message.data, .close); + return error.EndOfStream; + }, + else => unreachable, + } + } + } +}; + +/// Concurrent relay between an accepted WebSocket and any raw 9P byte stream. +/// Suitable for socket, pipe, serial, or firmware-provided std.Io adapters. +/// Both frame buffers must be disjoint and at least frame_limit bytes long. +/// Each Bridge value runs once. All streams and buffers remain caller-owned. Cancellation stops both pumps +/// before returning, including when the other peer is blocked waiting for I/O. +pub const Bridge = struct { + socket: *WebSocket, + upstream_reader: *Io.Reader, + upstream_writer: *Io.Writer, + request_buffer: []u8, + reply_buffer: []u8, + frame_limit: u32, + mutex: Io.Mutex = .init, + done: Io.Event = .unset, + errors: [2]?anyerror = .{ null, null }, + finished: std.atomic.Value(u8) = .init(2), + + pub fn run(b: *Bridge, io: Io, timeout: Io.Timeout) !void { + std.debug.assert(b.socket.role == .server); + std.debug.assert(b.frame_limit >= 24); + std.debug.assert(b.request_buffer.len >= b.frame_limit and b.reply_buffer.len >= b.frame_limit); + var group: Io.Group = .init; + defer group.cancel(io); + try group.concurrent(io, pump, .{ b, io, 0 }); + try group.concurrent(io, pump, .{ b, io, 1 }); + try b.done.waitTimeout(io, timeout); + group.cancel(io); + if (b.errors[b.finished.load(.acquire)]) |err| return err; + } + + fn pump(b: *Bridge, io: Io, direction: u1) void { + defer { + if (b.finished.cmpxchgStrong(2, direction, .release, .monotonic) == null) b.done.set(io); + } + if (direction == 0) b.requests(io) catch |err| { + b.errors[0] = err; + } else b.replies(io) catch |err| { + b.errors[1] = err; + }; + } + fn send(b: *Bridge, io: Io, bytes: []const u8, opcode: Opcode) !void { + try b.mutex.lock(io); + defer b.mutex.unlock(io); + try b.socket.send(bytes, opcode, null); + } + fn requests(b: *Bridge, io: Io) !void { + while (true) { + const message = try b.socket.receive(b.request_buffer[0..b.frame_limit]); + switch (message.opcode) { + .ping => try b.send(io, message.data, .pong), + .pong => {}, + .close => { + try b.send(io, message.data, .close); + return; + }, + .binary => { + const decoded = try wire.decode(message.data); + if (!wire.isT(decoded.msg.msgType())) return error.ExpectedRequest; + if (decoded.msg == .tversion and decoded.msg.tversion.msize > b.frame_limit) return error.MessageTooLarge; + try @import("transport.zig").writeFrame(b.upstream_writer, message.data, b.frame_limit); + try b.upstream_writer.flush(); + }, + else => unreachable, + } + } + } + fn replies(b: *Bridge, io: Io) !void { + while (true) { + const frame = try @import("transport.zig").readFrame(b.upstream_reader, b.reply_buffer, b.frame_limit); + const decoded = try wire.decode(frame); + if (wire.isT(decoded.msg.msgType())) return error.ExpectedReply; + try b.send(io, frame, .binary); + } + } +}; |
