summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/http.zig339
-rw-r--r--src/root.zig1
2 files changed, 340 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);
+ }
+ }
+};
diff --git a/src/root.zig b/src/root.zig
index 47303be..dadeb6d 100644
--- a/src/root.zig
+++ b/src/root.zig
@@ -46,6 +46,7 @@ pub const orclose: u8 = 64;
test {
@import("std").testing.refAllDecls(@This());
}
+pub const http = @import("http.zig");
pub const transport = @import("transport.zig");
pub const Quic = @import("quic.zig").Quic;
test {