diff options
| author | Gabriel Schneider <[email protected]> | 2026-09-14 14:10:28 -0300 |
|---|---|---|
| committer | Gabriel Schneider <[email protected]> | 2026-09-14 14:20:25 -0300 |
| commit | 5f24c4a2284a0af84fb9d116f7a39f6a58e42ae9 (patch) | |
| tree | 8679763439492361fe99ea01f5190e5204fa590c /src | |
| download | cloud9-5f24c4a2284a0af84fb9d116f7a39f6a58e42ae9.tar.gz cloud9-5f24c4a2284a0af84fb9d116f7a39f6a58e42ae9.zip | |
Implement base 9P2000 sessions, shared transports, and conformance probes
Diffstat (limited to 'src')
| -rw-r--r-- | src/Server.zig | 194 | ||||
| -rw-r--r-- | src/client.zig | 392 | ||||
| -rw-r--r-- | src/quic.zig | 302 | ||||
| -rw-r--r-- | src/root.zig | 53 | ||||
| -rw-r--r-- | src/session_test.zig | 222 | ||||
| -rw-r--r-- | src/transport.zig | 207 | ||||
| -rw-r--r-- | src/wire.zig | 1153 |
7 files changed, 2523 insertions, 0 deletions
diff --git a/src/Server.zig b/src/Server.zig new file mode 100644 index 0000000..a652f99 --- /dev/null +++ b/src/Server.zig @@ -0,0 +1,194 @@ +//! A bounded, caller-driven 9P2000 server connection. +//! The backend owns fids, authentication, permissions, and filesystem operations. +//! receive() borrows one input frame until release(). Pending operations may outlive +//! that frame only if the backend copies their strings/data. reply() copies output. +const std = @import("std"); +const assert = std.debug.assert; +const wire = @import("wire.zig"); +const Server = @This(); + +in: []u8, +out: []u8, +in_len: usize = 0, +frame: u32 = 0, +out_len: usize = 0, +out_off: usize = 0, +msize: u32 = 0, +dead: bool = false, +pending: [65]Pending = @splat(.{}), +versioning: bool = false, + +const Pending = struct { + request: ?wire.Type = null, + tag: u16 = 0, + count: u32 = 0, + oldtag: u16 = wire.notag, +}; + +pub const Error = wire.Error || error{ Protocol, NoTags, UnknownTag, WrongReply, TooLarge }; +pub const Options = struct { in: []u8, out: []u8 }; +/// Local resource floor, not a minimum imposed by the 9P specification. +pub const msize_min: u32 = 24; + +pub fn init(options: Options) Server { + assert(options.in.len >= msize_min); + assert(options.out.len >= msize_min); + return .{ .in = options.in, .out = options.out }; +} + +pub fn push(s: *Server, bytes: []const u8) usize { + if (s.dead) return 0; + const n = @min(bytes.len, s.in.len - s.in_len); + @memcpy(s.in[s.in_len..][0..n], bytes[0..n]); + s.in_len += n; + return n; +} + +pub fn output(s: *const Server) []const u8 { + return s.out[s.out_off..s.out_len]; +} + +pub fn wrote(s: *Server, n: usize) void { + assert(n <= s.output().len); + s.out_off += n; + if (s.out_off == s.out_len) { + s.out_off = 0; + s.out_len = 0; + } +} + +pub fn hasRoom(s: *Server) bool { + if (s.out_off != 0) { + const n = s.output().len; + std.mem.copyForwards(u8, s.out[0..n], s.output()); + s.out_off = 0; + s.out_len = n; + } + return s.out.len - s.out_len >= @max(s.msize, msize_min); +} + +/// A null result means more input or output drainage is needed. +/// Any receive error is terminal: the stream cannot be safely resynchronized. +pub fn receive(s: *Server) Error!?wire.Decoded { + assert(s.frame == 0); + if (s.dead) return null; + errdefer s.dead = true; + const len = wire.frameLen(s.in[0..s.in_len]) orelse return null; + if (len < wire.header_len or len > s.in.len) return error.Protocol; + if (len > s.in_len) return null; + const got = try wire.decode(s.in[0..len]); + const kind = got.msg.msgType(); + if (!wire.isT(kind)) return error.Protocol; + if (!s.hasRoom()) return null; + if (kind == .tversion) { + if (got.tag != wire.notag) return error.Protocol; + // Finish any partially transmitted response before starting a new session. + if (s.output().len != 0) return null; + s.pending = @splat(.{}); + s.versioning = true; + s.msize = 0; + } else { + if (s.msize == 0 or s.versioning) return error.Protocol; + if (len > s.msize or got.tag == wire.notag) return error.Protocol; + if (s.find(got.tag) != null) return error.Protocol; + const slot = s.free(kind) orelse return error.NoTags; + slot.* = .{ .request = kind, .tag = got.tag, .count = switch (got.msg) { + .tread => |m| m.count, + .twrite => |m| @intCast(m.data.len), + .twalk => |m| m.nwname, + else => 0, + }, .oldtag = if (got.msg == .tflush) got.msg.tflush.oldtag else wire.notag }; + } + s.frame = len; + return got; +} + +pub fn release(s: *Server) void { + assert(s.frame != 0 and s.frame <= s.in_len); + const n = s.in_len - s.frame; + std.mem.copyForwards(u8, s.in[0..n], s.in[s.frame..s.in_len]); + s.in_len = n; + s.frame = 0; +} + +/// Call only for a received Tversion, after aborting backend work and releasing fids. +/// Suffixes may fall back to base 9P2000; arbitrary strings beginning with 9P may not. +pub fn negotiate(s: *Server, want: u32, version: []const u8) Error!void { + assert(s.versioning); + const size: u32 = @intCast(@min(want, s.in.len, s.out.len, std.math.maxInt(u32))); + if (size < msize_min) return error.TooLarge; + const base = std.mem.sliceTo(version, '.'); + const known = std.mem.eql(u8, base, "9P2000"); + try s.append(wire.notag, .{ .rversion = .{ + .msize = size, + .version = if (known) "9P2000" else "unknown", + } }, size); + s.msize = if (known) size else 0; + s.versioning = false; +} + +/// An Rflush is a backend promise: no further response for oldtag will be sent. +/// A canceled backend operation must be retired before its tag can be reused. +pub fn reply(s: *Server, tag: u16, msg: wire.Msg) Error!void { + if (s.dead) return error.Protocol; + const slot = s.find(tag) orelse return error.UnknownTag; + const request = slot.request.?; + const kind = msg.msgType(); + if (request == .tflush and kind != .rflush) return error.WrongReply; + if (kind != .rerror and @intFromEnum(kind) != @intFromEnum(request) + 1) + return error.WrongReply; + switch (msg) { + .rread => |m| if (m.data.len > slot.count) return error.WrongReply, + .rwrite => |m| if (m.count > slot.count) return error.WrongReply, + .rwalk => |m| { + if (m.nwqid > slot.count) return error.WrongReply; + if (m.nwqid == 0 and slot.count != 0) return error.WrongReply; + }, + else => {}, + } + var response = msg; + if (response == .rerror) { + const cap = @min(s.msize - wire.header_len - 2, std.math.maxInt(u16)); + response.rerror.ename = response.rerror.ename[0..@min(response.rerror.ename.len, cap)]; + } + try s.append(tag, response, s.msize); + const oldtag = slot.oldtag; + slot.* = .{}; + if (kind == .rflush) { + if (s.find(oldtag)) |old| old.* = .{}; + } +} + +fn append(s: *Server, tag: u16, msg: wire.Msg, limit: u32) Error!void { + if (try wire.encodedLen(msg) > limit) return error.TooLarge; + _ = s.hasRoom(); + const bytes = try wire.encode(msg, tag, s.out[s.out_len..]); + s.out_len += bytes.len; +} + +fn find(s: *Server, tag: u16) ?*Pending { + for (&s.pending) |*slot| { + if (slot.request != null and slot.tag == tag) return slot; + } + return null; +} + +fn free(s: *Server, request: wire.Type) ?*Pending { + for (&s.pending, 0..) |*slot, i| { + // A full ordinary request window must still admit cancellation. + if (i == s.pending.len - 1 and request != .tflush) continue; + if (slot.request == null) return slot; + } + return null; +} + +pub fn hangup(s: *Server) void { + s.dead = true; + s.pending = @splat(.{}); + s.in_len = 0; + s.frame = 0; + s.out_len = 0; + s.out_off = 0; + s.msize = 0; + s.versioning = false; +} diff --git a/src/client.zig b/src/client.zig new file mode 100644 index 0000000..5b428dd --- /dev/null +++ b/src/client.zig @@ -0,0 +1,392 @@ +//! Caller-driven 9P2000 client. No allocator, sockets, threads, or filesystem policy. +//! Result slices remain valid until the next take() or hangup(). +const std = @import("std"); +const assert = std.debug.assert; +const wire = @import("wire.zig"); +const Msg = wire.Msg; +const Stat = wire.Stat; +const Qid = wire.Qid; +const notag = wire.notag; +const nofid = wire.nofid; +const max_welem = wire.max_welem; +const header_len = wire.header_len; +const isT = wire.isT; +const encode = wire.encode; +const decode = wire.decode; +const frameLen = wire.frameLen; +const Decoded = wire.Decoded; +const totalLen = wire.encodedLen; +pub const msize_min: u32 = 24; +pub const max_tags: usize = 16; + +const twrite_header: usize = header_len + 4 + 8 + 4; + +const rread_header: usize = header_len + 4; + +comptime { + assert(twrite_header == 23); + assert(rread_header == 11); + assert(max_tags <= notag); + assert(msize_min > rread_header); + assert(msize_min > twrite_header); +} + +pub const ClientError = error{ + NoTags, + NoSpace, + TooLarge, + Handshake, + Dead, + BadRequest, +}; + +pub const Client = struct { + in: []u8, + out: []u8, + + in_len: usize = 0, + frame: u32 = 0, + out_len: usize = 0, + out_off: usize = 0, + + msize: u32 = 0, + asked: u32 = 0, + versioning: bool = false, + dead: bool = false, + + tags: [max_tags + 1]Slot = @splat(.{}), + + const Slot = struct { + op: ?Op = null, + count: u32 = 0, + oldtag: u16 = notag, + completed: bool = false, + }; + + pub const Op = enum { version, auth, attach, flush, walk, open, create, read, write, clunk, remove, stat, wstat }; + + pub const Request = union(Op) { + version: struct { msize: u32 = 0 }, + auth: struct { afid: u32, uname: []const u8, aname: []const u8 = "" }, + attach: struct { fid: u32, afid: u32 = nofid, uname: []const u8, aname: []const u8 = "" }, + flush: struct { oldtag: u16 }, + walk: struct { fid: u32, newfid: u32, names: []const []const u8 }, + open: struct { fid: u32, mode: u8 }, + create: struct { fid: u32, name: []const u8, perm: u32, mode: u8 }, + read: struct { fid: u32, offset: u64, count: u32 }, + write: struct { fid: u32, offset: u64, data: []const u8 }, + clunk: struct { fid: u32 }, + remove: struct { fid: u32 }, + stat: struct { fid: u32 }, + wstat: struct { fid: u32, stat: Stat }, + }; + + pub const Result = union(enum) { + fail: []const u8, + version: struct { msize: u32, version: []const u8 }, + auth: Qid, + flush: void, + create: struct { qid: Qid, iounit: u32 }, + remove: void, + wstat: void, + attach: Qid, + walk: struct { nwqid: u16, wqid: [max_welem]Qid }, + open: struct { qid: Qid, iounit: u32 }, + read: []const u8, + write: u32, + clunk: void, + stat: Stat, + }; + + pub const Done = struct { + tag: u16, + op: Op, + result: Result, + }; + + pub const Options = struct { + in: []u8, + out: []u8, + }; + + pub fn init(opts: Options) Client { + assert(opts.in.len >= msize_min); + assert(opts.out.len >= msize_min); + return .{ .in = opts.in, .out = opts.out }; + } + + pub fn hangup(c: *Client) void { + c.dead = true; + c.tags = @splat(.{}); + c.versioning = false; + c.msize = 0; + c.asked = 0; + c.in_len = 0; + c.frame = 0; + c.out_len = 0; + c.out_off = 0; + } + + pub fn push(c: *Client, bytes: []const u8) usize { + if (c.dead) return 0; + const n = @min(bytes.len, c.in.len - c.in_len); + @memcpy(c.in[c.in_len..][0..n], bytes[0..n]); + c.in_len += n; + return n; + } + + pub fn output(c: *const Client) []const u8 { + return c.out[c.out_off..c.out_len]; + } + + pub fn wrote(c: *Client, n: usize) void { + assert(n <= c.out_len - c.out_off); + c.out_off += n; + if (c.out_off == c.out_len) { + c.out_off = 0; + c.out_len = 0; + } + } + + fn compact(c: *Client) void { + assert(c.out_off <= c.out_len); + const n = c.out_len - c.out_off; + std.mem.copyForwards(u8, c.out[0..n], c.out[c.out_off..c.out_len]); + c.out_off = 0; + c.out_len = n; + } + + fn dropFrame(c: *Client) void { + assert(c.frame != 0); + assert(c.frame <= c.in_len); + const n = c.frame; + std.mem.copyForwards(u8, c.in[0 .. c.in_len - n], c.in[n..c.in_len]); + c.in_len -= n; + c.frame = 0; + } + + pub fn maxRead(c: *const Client) u32 { + if (c.msize == 0) return 0; + return c.msize - @as(u32, @intCast(rread_header)); + } + + pub fn maxWrite(c: *const Client) u32 { + if (c.msize == 0) return 0; + return c.msize - @as(u32, @intCast(twrite_header)); + } + + pub fn pending(c: *const Client) usize { + var n: usize = @intFromBool(c.versioning); + for (c.tags) |t| n += @intFromBool(t.op != null); + return n; + } + + pub fn submit(c: *Client, req: Request) ClientError!u16 { + if (c.dead) return error.Dead; + if (req == .version) return c.beginVersion(req.version.msize); + if (c.msize == 0 or c.versioning) return error.Handshake; + + const msg: Msg = switch (req) { + .version => unreachable, // handled above + .auth => |m| blk: { + if (m.afid == nofid) return error.BadRequest; + break :blk .{ .tauth = .{ .afid = m.afid, .uname = m.uname, .aname = m.aname } }; + }, + .flush => |m| blk: { + if (m.oldtag == notag) return error.BadRequest; + break :blk .{ .tflush = .{ .oldtag = m.oldtag } }; + }, + .create => |m| blk: { + if (m.fid == nofid or m.name.len == 0) return error.BadRequest; + if (std.mem.indexOfAny(u8, m.name, "/\x00") != null) return error.BadRequest; + if (std.mem.eql(u8, m.name, ".") or std.mem.eql(u8, m.name, "..")) return error.BadRequest; + break :blk .{ .tcreate = .{ .fid = m.fid, .name = m.name, .perm = m.perm, .mode = m.mode } }; + }, + .remove => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + break :blk .{ .tremove = .{ .fid = m.fid } }; + }, + .wstat => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + break :blk .{ .twstat = .{ .fid = m.fid, .stat = m.stat } }; + }, + .attach => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + break :blk .{ .tattach = .{ + .fid = m.fid, + .afid = m.afid, + .uname = m.uname, + .aname = m.aname, + } }; + }, + .walk => |m| blk: { + if (m.fid == nofid or m.newfid == nofid) return error.BadRequest; + if (m.names.len > max_welem) return error.BadRequest; + var w: [max_welem][]const u8 = @splat(""); + for (m.names, 0..) |n, i| { + if (n.len == 0) return error.BadRequest; + if (std.mem.indexOfAny(u8, n, "/\x00") != null) return error.BadRequest; + w[i] = n; + } + break :blk .{ .twalk = .{ + .fid = m.fid, + .newfid = m.newfid, + .nwname = @intCast(m.names.len), + .wname = w, + } }; + }, + .open => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + break :blk .{ .topen = .{ .fid = m.fid, .mode = m.mode } }; + }, + .read => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + if (m.count > c.maxRead()) return error.TooLarge; + break :blk .{ .tread = .{ .fid = m.fid, .offset = m.offset, .count = m.count } }; + }, + .write => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + break :blk .{ .twrite = .{ .fid = m.fid, .offset = m.offset, .data = m.data } }; + }, + .clunk => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + break :blk .{ .tclunk = .{ .fid = m.fid } }; + }, + .stat => |m| blk: { + if (m.fid == nofid) return error.BadRequest; + break :blk .{ .tstat = .{ .fid = m.fid } }; + }, + }; + const need = totalLen(msg) catch return error.TooLarge; + if (need > c.msize) return error.TooLarge; + + const op = std.meta.activeTag(req); + const tag = c.claim(op, if (req == .flush) req.flush.oldtag else notag) orelse return error.NoTags; + errdefer c.tags[tag] = .{}; + try c.emit(tag, msg); + switch (req) { + .read => |m| c.tags[tag].count = m.count, + .write => |m| c.tags[tag].count = @intCast(m.data.len), + .walk => |m| c.tags[tag].count = @intCast(m.names.len), + .flush => |m| { + c.tags[tag].oldtag = m.oldtag; + }, + else => {}, + } + return tag; + } + + fn beginVersion(c: *Client, want: u32) ClientError!u16 { + if (c.pending() != 0) return error.Handshake; + const cap: u32 = @intCast(@min(c.in.len, c.out.len, std.math.maxInt(u32))); + const m = @min(if (want == 0) cap else want, cap); + if (m < msize_min) return error.BadRequest; + try c.emit(notag, .{ .tversion = .{ .msize = m, .version = "9P2000" } }); + c.msize = 0; + c.asked = m; + c.versioning = true; + return notag; + } + + fn reserved(c: *const Client, tag: u16) bool { + for (c.tags) |slot| { + if (slot.op == .flush and !slot.completed and slot.oldtag == tag) return true; + } + return false; + } + + fn claim(c: *Client, op: Op, avoid: u16) ?u16 { + for (&c.tags, 0..) |*t, i| { + // Keep one tag available to cancel a fully occupied request window. + if (i >= max_tags and op != .flush) continue; + if (t.op != null or i == avoid or c.reserved(@intCast(i))) continue; + t.* = .{ .op = op }; + return @intCast(i); + } + return null; + } + + fn emit(c: *Client, tag: u16, msg: Msg) ClientError!void { + if (c.out_off != 0) c.compact(); + const bytes = encode(msg, tag, c.out[c.out_len..]) catch return error.NoSpace; + c.out_len += bytes.len; + } + + pub fn take(c: *Client) ?Done { + if (c.frame != 0) c.dropFrame(); + if (c.dead) return null; + const len = frameLen(c.in[0..c.in_len]) orelse return null; + if (len < header_len or len > c.in.len) return c.die(); + if (c.msize != 0 and len > c.msize) return c.die(); + if (len > c.in_len) return null; + c.frame = len; + const got = decode(c.in[0..len]) catch return c.die(); + return c.consume(got); + } + + fn die(c: *Client) ?Done { + c.dead = true; + return null; + } + + fn consume(c: *Client, got: Decoded) ?Done { + if (isT(got.msg.msgType())) return c.die(); + if (got.msg == .rversion) return c.version(got); + if (c.versioning or c.msize == 0) return c.die(); + if (got.tag >= c.tags.len) return c.die(); + const slot = &c.tags[got.tag]; + const op = slot.op orelse return c.die(); + if (slot.completed) return c.die(); + const result: Result = switch (got.msg) { + .rerror => |m| if (op == .flush) return c.die() else .{ .fail = m.ename }, + .rauth => |m| if (op != .auth) return c.die() else .{ .auth = m.aqid }, + .rcreate => |m| if (op != .create) return c.die() else .{ .create = .{ .qid = m.qid, .iounit = m.iounit } }, + .rremove => if (op != .remove) return c.die() else .remove, + .rwstat => if (op != .wstat) return c.die() else .wstat, + .rflush => blk: { + if (op != .flush) return c.die(); + if (slot.oldtag < c.tags.len) c.tags[slot.oldtag] = .{}; + break :blk .flush; + }, + .rattach => |m| if (op != .attach) return c.die() else .{ .attach = m.qid }, + .rwalk => |m| if (op != .walk or m.nwqid > slot.count or (m.nwqid == 0 and slot.count != 0)) return c.die() else .{ + .walk = .{ .nwqid = m.nwqid, .wqid = m.wqid }, + }, + .ropen => |m| if (op != .open) return c.die() else .{ + .open = .{ .qid = m.qid, .iounit = m.iounit }, + }, + .rread => |m| blk: { + if (op != .read) return c.die(); + if (m.data.len > slot.count) return c.die(); + break :blk .{ .read = m.data }; + }, + .rwrite => |m| if (op != .write or m.count > slot.count) return c.die() else .{ .write = m.count }, + .rclunk => if (op != .clunk) return c.die() else .clunk, + .rstat => |m| if (op != .stat) return c.die() else .{ .stat = m.stat }, + else => return c.die(), + }; + slot.completed = true; + // Multiple flushes can reserve the same tag; flushing a flush can release + // a reservation on a completed original operation. + for (&c.tags, 0..) |*candidate, tag| { + if (candidate.completed and !c.reserved(@intCast(tag))) candidate.* = .{}; + } + return .{ .tag = got.tag, .op = op, .result = result }; + } + + fn version(c: *Client, got: Decoded) ?Done { + if (!c.versioning) return c.die(); + if (got.tag != notag) return c.die(); + const m = got.msg.rversion; + if (m.msize > c.asked or m.msize < msize_min) return c.die(); + c.versioning = false; + if (std.mem.eql(u8, m.version, "9P2000")) { + c.msize = m.msize; + } else if (!std.mem.eql(u8, m.version, "unknown")) { + return c.die(); + } + return .{ .tag = notag, .op = .version, .result = .{ + .version = .{ .msize = m.msize, .version = m.version }, + } }; + } +}; diff --git a/src/quic.zig b/src/quic.zig new file mode 100644 index 0000000..c78a79d --- /dev/null +++ b/src/quic.zig @@ -0,0 +1,302 @@ +const std = @import("std"); +const libc = std.c; + +/// QUIC transport with caller-selected ALPN. OpenSSL bindings are injected so +/// applications control linkage. Certificates are ephemeral and peers unauthenticated. +pub fn Quic(comptime ssl: type, comptime protocol: []const u8) type { + if (protocol.len == 0 or protocol.len > 255) @compileError("invalid ALPN length"); + return struct { + comptime { + if (ssl.OPENSSL_VERSION_NUMBER < 0x30600000) + @compileError("9P over QUIC requires OpenSSL 3.6 or newer"); + } + + pub const alpn = protocol; + pub const Error = error{ Tls, Socket, SocketFlags, SocketOption, Bind, Address, Closed, InvalidWrite }; + + pub const Listener = struct { + fd: c_int, + handle: *ssl.SSL, + address: std.Io.net.IpAddress, + + pub fn init(address: std.Io.net.IpAddress) Error!Listener { + ssl.ERR_clear_error(); + const ctx = ssl.SSL_CTX_new(ssl.OSSL_QUIC_server_method()) orelse return error.Tls; + defer ssl.SSL_CTX_free(ctx); + const key = ssl.EVP_PKEY_Q_keygen(null, null, "EC", @as([*:0]const u8, "prime256v1")) orelse return error.Tls; + defer ssl.EVP_PKEY_free(key); + const cert = ssl.X509_new() orelse return error.Tls; + defer ssl.X509_free(cert); + if (ssl.X509_set_version(cert, 2) != 1 or + ssl.ASN1_INTEGER_set(ssl.X509_get_serialNumber(cert), 1) != 1 or + ssl.X509_gmtime_adj(ssl.X509_getm_notBefore(cert), -60) == null or + ssl.X509_gmtime_adj(ssl.X509_getm_notAfter(cert), 365 * 24 * 60 * 60) == null or + ssl.X509_set_pubkey(cert, key) != 1) return error.Tls; + const name = ssl.X509_get_subject_name(cert) orelse return error.Tls; + if (ssl.X509_NAME_add_entry_by_txt(name, "CN", ssl.MBSTRING_ASC, protocol.ptr, @intCast(protocol.len), -1, 0) != 1 or + ssl.X509_set_issuer_name(cert, name) != 1 or + ssl.X509_sign(cert, key, ssl.EVP_sha256()) <= 0 or + ssl.SSL_CTX_use_certificate(ctx, cert) != 1 or + ssl.SSL_CTX_use_PrivateKey(ctx, key) != 1) return error.Tls; + ssl.SSL_CTX_set_verify(ctx, ssl.SSL_VERIFY_NONE, null); + ssl.SSL_CTX_set_alpn_select_cb(ctx, selectAlpn, null); + var addr: libc.sockaddr.storage = undefined; + const addr_len = sockaddr(address, &addr); + const fd = try udp(addr.family); + errdefer _ = libc.close(fd); + if (libc.bind(fd, @ptrCast(&addr), addr_len) != 0) return error.Bind; + var actual_len: libc.socklen_t = @sizeOf(@TypeOf(addr)); + if (libc.getsockname(fd, @ptrCast(&addr), &actual_len) != 0) return error.Address; + var actual = address; + actual.setPort(switch (address) { + .ip4 => std.mem.bigToNative(u16, @as(*const libc.sockaddr.in, @ptrCast(&addr)).port), + .ip6 => std.mem.bigToNative(u16, @as(*const libc.sockaddr.in6, @ptrCast(&addr)).port), + }); + const handle = ssl.SSL_new_listener(ctx, 0) orelse return error.Tls; + errdefer ssl.SSL_free(handle); + if (ssl.SSL_set_fd(handle, fd) != 1 or ssl.SSL_set_blocking_mode(handle, 0) != 1 or + ssl.SSL_listen(handle) != 1) return error.Tls; + return .{ .fd = fd, .handle = handle, .address = actual }; + } + + pub fn accept(l: *Listener) Error!?Connection { + ssl.ERR_clear_error(); + const handle = ssl.SSL_accept_connection(l.handle, ssl.SSL_ACCEPT_CONNECTION_NO_BLOCK) orelse { + if (ssl.ERR_peek_error() != 0) return error.Tls; + return null; + }; + errdefer ssl.SSL_free(handle); + if (ssl.SSL_set_default_stream_mode(handle, ssl.SSL_DEFAULT_STREAM_MODE_NONE) != 1 or + ssl.SSL_set_blocking_mode(handle, 0) != 1) return error.Tls; + return .{ .handle = handle }; + } + + pub fn events(l: *Listener) Error!void { + ssl.ERR_clear_error(); + if (ssl.SSL_handle_events(l.handle) != 1) return error.Tls; + } + + pub fn poll(l: *const Listener) libc.pollfd { + return pollFd(l.handle, l.fd); + } + + pub fn nextDue(l: *const Listener) ?i32 { + return due(l.handle); + } + + // Accepted connections must be released before the shared UDP socket. + pub fn deinit(l: *Listener) void { + ssl.SSL_free(l.handle); + _ = libc.close(l.fd); + l.* = undefined; + } + }; + + pub const Connection = struct { + handle: *ssl.SSL, + stream: ?*ssl.SSL = null, + fd: c_int = -1, + pending_write_len: usize = 0, + + pub fn dial(address: std.Io.net.IpAddress) Error!Connection { + ssl.ERR_clear_error(); + const ctx = ssl.SSL_CTX_new(ssl.OSSL_QUIC_client_method()) orelse return error.Tls; + defer ssl.SSL_CTX_free(ctx); + ssl.SSL_CTX_set_verify(ctx, ssl.SSL_VERIFY_NONE, null); + const fd = try udp(if (address == .ip4) libc.AF.INET else libc.AF.INET6); + errdefer _ = libc.close(fd); + const handle = ssl.SSL_new(ctx) orelse return error.Tls; + errdefer ssl.SSL_free(handle); + if (ssl.SSL_set_fd(handle, fd) != 1 or ssl.SSL_set_blocking_mode(handle, 0) != 1 or + ssl.SSL_set_default_stream_mode(handle, ssl.SSL_DEFAULT_STREAM_MODE_NONE) != 1) return error.Tls; + const protocols = [_]u8{alpn.len} ++ protocol[0..protocol.len].*; + if (ssl.SSL_set_alpn_protos(handle, &protocols, protocols.len) != 0) return error.Tls; + const peer = ssl.BIO_ADDR_new() orelse return error.Tls; + defer ssl.BIO_ADDR_free(peer); + const made = switch (address) { + .ip4 => |ip| ssl.BIO_ADDR_rawmake(peer, libc.AF.INET, &ip.bytes, ip.bytes.len, std.mem.nativeToBig(u16, ip.port)), + .ip6 => |ip| ssl.BIO_ADDR_rawmake(peer, libc.AF.INET6, &ip.bytes, ip.bytes.len, std.mem.nativeToBig(u16, ip.port)), + }; + if (made != 1 or ssl.SSL_set1_initial_peer_addr(handle, peer) != 1) return error.Tls; + return .{ .handle = handle, .fd = fd }; + } + + pub fn handshake(c: *Connection) Error!bool { + ssl.ERR_clear_error(); + var close_info: ssl.SSL_CONN_CLOSE_INFO = undefined; + if (ssl.SSL_get_conn_close_info(c.handle, &close_info, @sizeOf(@TypeOf(close_info))) == 1) + return error.Closed; + if (ssl.SSL_is_init_finished(c.handle) == 1) return true; + const rc = if (c.fd >= 0) ssl.SSL_connect(c.handle) else ssl.SSL_accept(c.handle); + if (rc == 1) return true; + try retry(c.handle, rc); + return false; + } + + fn ready(c: *Connection) Error!bool { + if (!try c.handshake()) return false; + if (c.stream != null) return true; + ssl.ERR_clear_error(); + const stream = if (c.fd >= 0) + ssl.SSL_new_stream(c.handle, ssl.SSL_STREAM_FLAG_NO_BLOCK) + else + ssl.SSL_accept_stream(c.handle, ssl.SSL_ACCEPT_STREAM_NO_BLOCK); + if (stream == null) { + if (ssl.ERR_peek_error() != 0) return error.Tls; + return false; + } + errdefer ssl.SSL_free(stream); + if (ssl.SSL_set_blocking_mode(stream, 0) != 1 or + ssl.SSL_get_stream_id(stream) != 0 or + ssl.SSL_set_incoming_stream_policy(c.handle, ssl.SSL_INCOMING_STREAM_POLICY_REJECT, 0) != 1) + return error.Tls; + _ = ssl.SSL_set_mode(stream, ssl.SSL_MODE_ACCEPT_MOVING_WRITE_BUFFER); + c.stream = stream; + return true; + } + + pub fn read(c: *Connection, bytes: []u8) Error!?usize { + if (!try c.ready()) return null; + ssl.ERR_clear_error(); + var len: usize = 0; + const rc = ssl.SSL_read_ex(c.stream, bytes.ptr, bytes.len, &len); + if (rc == 1) return len; + if (ssl.SSL_get_error(c.stream, rc) == ssl.SSL_ERROR_ZERO_RETURN) return 0; + try retry(c.stream.?, rc); + return null; + } + + pub fn pending(c: *const Connection) bool { + var close_info: ssl.SSL_CONN_CLOSE_INFO = undefined; + if (ssl.SSL_get_conn_close_info(c.handle, &close_info, @sizeOf(@TypeOf(close_info))) == 1) return true; + if (c.stream) |stream| { + var item: ssl.SSL_POLL_ITEM = .{ .desc = ssl.SSL_as_poll_descriptor(stream), .events = ssl.SSL_POLL_EVENT_RE, .revents = 0 }; + const timeout: ssl.struct_timeval = .{ .tv_sec = 0, .tv_usec = 0 }; + if (ssl.SSL_poll(&item, 1, @sizeOf(@TypeOf(item)), &timeout, ssl.SSL_POLL_FLAG_NO_HANDLE_EVENTS, null) != 1) return true; + return item.revents != 0; + } + return ssl.SSL_get_accept_stream_queue_len(c.handle) != 0; + } + + pub fn write(c: *Connection, bytes: []const u8) Error!usize { + if (!try c.ready()) return 0; + if (bytes.len < c.pending_write_len) return error.InvalidWrite; + const requested = if (c.pending_write_len != 0) c.pending_write_len else bytes.len; + if (requested == 0) return 0; + ssl.ERR_clear_error(); + var len: usize = 0; + const rc = ssl.SSL_write_ex(c.stream, bytes.ptr, requested, &len); + if (rc == 1) { + c.pending_write_len = 0; + return len; + } + try retry(c.stream.?, rc); + c.pending_write_len = requested; + return 0; + } + + pub fn conclude(c: *Connection) Error!void { + if (c.pending_write_len != 0) return error.InvalidWrite; + if (!try c.ready()) return error.Closed; + ssl.ERR_clear_error(); + if (ssl.SSL_stream_conclude(c.stream, 0) != 1) return error.Tls; + } + + pub fn events(c: *Connection) Error!void { + if (c.fd < 0) return; + ssl.ERR_clear_error(); + if (ssl.SSL_handle_events(c.handle) != 1) return error.Tls; + } + + pub fn poll(c: *const Connection) ?libc.pollfd { + return if (c.fd >= 0) pollFd(c.handle, c.fd) else null; + } + + pub fn nextDue(c: *const Connection) ?i32 { + return if (c.fd >= 0) due(c.handle) else null; + } + + pub fn deinit(c: *Connection) void { + ssl.ERR_clear_error(); + _ = ssl.SSL_shutdown_ex(c.handle, ssl.SSL_SHUTDOWN_FLAG_RAPID | ssl.SSL_SHUTDOWN_FLAG_NO_STREAM_FLUSH | ssl.SSL_SHUTDOWN_FLAG_NO_BLOCK, null, 0); + ssl.SSL_free(c.stream); + ssl.SSL_free(c.handle); + if (c.fd >= 0) _ = libc.close(c.fd); + c.* = undefined; + } + }; + + fn udp(family: u16) Error!c_int { + const fd = libc.socket(family, libc.SOCK.DGRAM, 0); + if (fd < 0) return error.Socket; + errdefer _ = libc.close(fd); + const flags = libc.fcntl(fd, libc.F.GETFL, @as(c_int, 0)); + if (flags < 0) return error.SocketFlags; + var options: libc.O = @bitCast(@as(u32, @bitCast(flags))); + options.NONBLOCK = true; + if (libc.fcntl(fd, libc.F.SETFL, @as(c_int, @bitCast(@as(u32, @bitCast(options))))) != 0 or + libc.fcntl(fd, libc.F.SETFD, @as(c_int, libc.FD_CLOEXEC)) != 0) return error.SocketFlags; + if (family == libc.AF.INET6) { + const enabled: c_int = 1; + const v6only = if (@import("builtin").os.tag.isDarwin()) 27 else libc.IPV6.V6ONLY; + if (libc.setsockopt(fd, libc.IPPROTO.IPV6, v6only, &enabled, @sizeOf(c_int)) != 0) return error.SocketOption; + } + return fd; + } + + fn sockaddr(address: std.Io.net.IpAddress, out: *libc.sockaddr.storage) libc.socklen_t { + switch (address) { + .ip4 => |ip| { + const addr: *libc.sockaddr.in = @ptrCast(out); + addr.* = .{ .port = std.mem.nativeToBig(u16, ip.port), .addr = @bitCast(ip.bytes) }; + return @sizeOf(libc.sockaddr.in); + }, + .ip6 => |ip| { + const addr: *libc.sockaddr.in6 = @ptrCast(out); + addr.* = .{ .port = std.mem.nativeToBig(u16, ip.port), .addr = ip.bytes, .flowinfo = 0, .scope_id = 0 }; + return @sizeOf(libc.sockaddr.in6); + }, + } + } + + fn pollFd(handle: *ssl.SSL, fd: c_int) libc.pollfd { + var result: libc.pollfd = .{ .fd = fd, .events = 0, .revents = 0 }; + if (ssl.SSL_net_read_desired(handle) == 1) result.events |= libc.POLL.IN; + if (ssl.SSL_net_write_desired(handle) == 1) result.events |= libc.POLL.OUT; + return result; + } + + fn due(handle: *ssl.SSL) ?i32 { + var tv: ssl.struct_timeval = undefined; + var infinite: c_int = undefined; + if (ssl.SSL_get_event_timeout(handle, &tv, &infinite) != 1) return 0; + if (infinite != 0) return null; + const ms = @as(i128, tv.tv_sec) * 1000 + @divFloor(@as(i128, tv.tv_usec) + 999, 1000); + return @intCast(std.math.clamp(ms, 0, std.math.maxInt(i32))); + } + + fn retry(handle: *ssl.SSL, rc: c_int) Error!void { + switch (ssl.SSL_get_error(handle, rc)) { + ssl.SSL_ERROR_WANT_READ, ssl.SSL_ERROR_WANT_WRITE => {}, + ssl.SSL_ERROR_ZERO_RETURN => return error.Closed, + else => return error.Tls, + } + } + + fn selectAlpn(_: ?*ssl.SSL, out: [*c][*c]const u8, outlen: [*c]u8, input: [*c]const u8, len: c_uint, _: ?*anyopaque) callconv(.c) c_int { + var offset: usize = 0; + while (offset < len) { + const size = input[offset]; + offset += 1; + if (size > len - offset) return ssl.SSL_TLSEXT_ERR_ALERT_FATAL; + if (std.mem.eql(u8, input[offset..][0..size], alpn)) { + out.* = input + offset; + outlen.* = size; + return ssl.SSL_TLSEXT_ERR_OK; + } + offset += size; + } + return ssl.SSL_TLSEXT_ERR_ALERT_FATAL; + } + }; +} diff --git a/src/root.zig b/src/root.zig new file mode 100644 index 0000000..47303be --- /dev/null +++ b/src/root.zig @@ -0,0 +1,53 @@ +//! Allocation-free 9P2000 protocol library. See docs/design.md for ownership contracts. +pub const wire = @import("wire.zig"); +pub const Error = wire.Error; +pub const Type = wire.Type; +pub const isT = wire.isT; +pub const header_len = wire.header_len; +pub const qid_len = wire.qid_len; +pub const stat_fixed = wire.stat_fixed; +pub const notag = wire.notag; +pub const nofid = wire.nofid; +pub const max_welem = wire.max_welem; +pub const iohdrsz = wire.iohdrsz; +pub const Qid = wire.Qid; +pub const Stat = wire.Stat; +pub const Msg = wire.Msg; +pub const Decoded = wire.Decoded; +pub const frameLen = wire.frameLen; +pub const encode = wire.encode; +pub const decode = wire.decode; +pub const encodedLen = wire.encodedLen; +pub const qtdir = wire.qtdir; +pub const qtappend = wire.qtappend; +pub const qtexcl = wire.qtexcl; +pub const qtmount = wire.qtmount; +pub const qtauth = wire.qtauth; +pub const qttmp = wire.qttmp; +pub const qtfile = wire.qtfile; +pub const dmdir = wire.dmdir; +pub const dmappend = wire.dmappend; +pub const dmexcl = wire.dmexcl; +pub const dmmount = wire.dmmount; +pub const dmauth = wire.dmauth; +pub const dmtmp = wire.dmtmp; +pub const dmperm = wire.dmperm; +pub const Client = @import("client.zig").Client; +pub const ClientError = @import("client.zig").ClientError; +pub const max_tags = @import("client.zig").max_tags; +pub const Server = @import("Server.zig"); +pub const oread: u8 = 0; +pub const owrite: u8 = 1; +pub const ordwr: u8 = 2; +pub const oexec: u8 = 3; +pub const otrunc: u8 = 16; +pub const ocexec: u8 = 32; +pub const orclose: u8 = 64; +test { + @import("std").testing.refAllDecls(@This()); +} +pub const transport = @import("transport.zig"); +pub const Quic = @import("quic.zig").Quic; +test { + _ = @import("session_test.zig"); +} diff --git a/src/session_test.zig b/src/session_test.zig new file mode 100644 index 0000000..8600d03 --- /dev/null +++ b/src/session_test.zig @@ -0,0 +1,222 @@ +const std = @import("std"); +const testing = std.testing; +const c9 = @import("root.zig"); +const Pair = struct { + client_in: [4096]u8 = undefined, + client_out: [4096]u8 = undefined, + server_in: [4096]u8 = undefined, + server_out: [8192]u8 = undefined, + client: c9.Client = undefined, + server: c9.Server = undefined, + + fn init(p: *Pair) !void { + p.client = .init(.{ .in = &p.client_in, .out = &p.client_out }); + p.server = .init(.{ .in = &p.server_in, .out = &p.server_out }); + _ = try p.client.submit(.{ .version = .{} }); + const request = try p.nextRequest(); + try p.server.negotiate(request.msg.tversion.msize, request.msg.tversion.version); + p.server.release(); + try testing.expectEqual(c9.Client.Op.version, (try p.nextResult()).op); + } + + fn nextRequest(p: *Pair) !c9.Decoded { + // Exercise every frame boundary through single-byte delivery. + while (p.client.output().len != 0) { + try testing.expectEqual(@as(usize, 1), p.server.push(p.client.output()[0..1])); + p.client.wrote(1); + if (try p.server.receive()) |request| return request; + } + return (try p.server.receive()) orelse error.NoRequest; + } + + fn nextResult(p: *Pair) !c9.Client.Done { + while (p.server.output().len != 0) { + try testing.expectEqual(@as(usize, 1), p.client.push(p.server.output()[0..1])); + p.server.wrote(1); + if (p.client.take()) |result| return result; + } + return p.client.take() orelse error.NoResult; + } + + fn exchange(p: *Pair, req: c9.Client.Request, response: c9.Msg) !void { + const tag = try p.client.submit(req); + const request = try p.nextRequest(); + try testing.expectEqual(tag, request.tag); + try p.server.reply(tag, response); + p.server.release(); + const result = try p.nextResult(); + try testing.expectEqual(std.meta.activeTag(req), result.op); + try testing.expectEqualStrings(@tagName(req), @tagName(result.result)); + } +}; + +const qid: c9.Qid = .{ .type = 0, .version = 1, .path = 2 }; +const stat: c9.Stat = .{ + .type = 0, + .dev = 0, + .qid = qid, + .mode = 0o600, + .atime = 0, + .mtime = 0, + .length = 0, + .name = "file", + .uid = "user", + .gid = "group", + .muid = "user", +}; + +test "all base 9P2000 client operations through fragmented server transport" { + var p: Pair = .{}; + try p.init(); + try p.exchange(.{ .auth = .{ .afid = 1, .uname = "user" } }, .{ .rauth = .{ .aqid = qid } }); + try p.exchange(.{ .attach = .{ .fid = 2, .afid = 1, .uname = "user" } }, .{ .rattach = .{ .qid = qid } }); + try p.exchange(.{ .walk = .{ .fid = 2, .newfid = 3, .names = &.{"file"} } }, .{ .rwalk = .{ .nwqid = 1, .wqid = @splat(qid) } }); + try p.exchange(.{ .open = .{ .fid = 3, .mode = 2 } }, .{ .ropen = .{ .qid = qid, .iounit = 0 } }); + try p.exchange(.{ .create = .{ .fid = 2, .name = "new", .perm = 0o600, .mode = 2 } }, .{ .rcreate = .{ .qid = qid, .iounit = 0 } }); + try p.exchange(.{ .read = .{ .fid = 3, .offset = 0, .count = 3 } }, .{ .rread = .{ .data = "abc" } }); + try p.exchange(.{ .write = .{ .fid = 3, .offset = 0, .data = "abc" } }, .{ .rwrite = .{ .count = 2 } }); + try p.exchange(.{ .stat = .{ .fid = 3 } }, .{ .rstat = .{ .stat = stat } }); + try p.exchange(.{ .wstat = .{ .fid = 3, .stat = stat } }, .rwstat); + try p.exchange(.{ .clunk = .{ .fid = 3 } }, .rclunk); + try p.exchange(.{ .remove = .{ .fid = 2 } }, .rremove); + try p.exchange(.{ .flush = .{ .oldtag = 42 } }, .rflush); +} + +test "flush holds oldtag until Rflush, including a completed original request" { + for ([_]bool{ false, true }) |complete| { + var p: Pair = .{}; + try p.init(); + const oldtag = try p.client.submit(.{ .read = .{ .fid = 1, .offset = 0, .count = 1 } }); + _ = try p.nextRequest(); + p.server.release(); + const flush = try p.client.submit(.{ .flush = .{ .oldtag = oldtag } }); + _ = try p.nextRequest(); + p.server.release(); + if (complete) { + try p.server.reply(oldtag, .{ .rread = .{ .data = "x" } }); + try testing.expectEqual(oldtag, (try p.nextResult()).tag); + try testing.expectEqual(@as(usize, 2), p.client.pending()); + } + const other = try p.client.submit(.{ .stat = .{ .fid = 1 } }); + try testing.expect(other != oldtag and other != flush); + try p.server.reply(flush, .rflush); + try testing.expectEqual(flush, (try p.nextResult()).tag); + try testing.expectEqual(@as(usize, 1), p.client.pending()); + try testing.expectError(error.UnknownTag, p.server.reply(oldtag, .{ .rread = .{ .data = "x" } })); + const reused = try p.client.submit(.{ .stat = .{ .fid = 1 } }); + try testing.expectEqual(oldtag, reused); + } +} + +test "response bounds and reply types are checked before output is changed" { + var p: Pair = .{}; + try p.init(); + const tag = try p.client.submit(.{ .write = .{ .fid = 1, .offset = 0, .data = "a" } }); + _ = try p.nextRequest(); + p.server.release(); + try testing.expectError(error.WrongReply, p.server.reply(tag, .{ .rwrite = .{ .count = 2 } })); + try testing.expectError(error.WrongReply, p.server.reply(tag, .rclunk)); + try testing.expectEqual(@as(usize, 0), p.server.output().len); + try p.server.reply(tag, .{ .rwrite = .{ .count = 1 } }); + _ = try p.nextResult(); +} + +test "stat outer length overflow is rejected without writing output" { + var name: [65535 - 48]u8 = @splat('a'); + var entry = stat; + entry.name = &name; + entry.uid = ""; + entry.gid = ""; + entry.muid = ""; + var buffer: [65550]u8 = @splat(0xaa); + try testing.expectError(error.Overlong, c9.encode(.{ .rstat = .{ .stat = entry } }, 1, &buffer)); + try testing.expect(std.mem.allEqual(u8, &buffer, 0xaa)); +} + +test "arbitrary input decoder round trip" { + try testing.fuzz({}, fuzzDecode, .{}); +} + +fn fuzzDecode(_: void, smith: *testing.Smith) !void { + var input_buffer: [65536]u8 = undefined; + const input = input_buffer[0..smith.slice(&input_buffer)]; + const decoded = c9.decode(input) catch return; + var buffer: [65536]u8 = undefined; + const encoded = try c9.encode(decoded.msg, decoded.tag, &buffer); + try testing.expectEqualSlices(u8, input, encoded); +} + +test "flush can cancel a full request window and reserves unused oldtags" { + var p: Pair = .{}; + try p.init(); + _ = try p.client.submit(.{ .flush = .{ .oldtag = 2 } }); + const next = try p.client.submit(.{ .stat = .{ .fid = 0 } }); + try testing.expect(next != 2); + const third = try p.client.submit(.{ .stat = .{ .fid = 0 } }); + try testing.expect(third != 2); + p.client.hangup(); + try p.init(); + for (0..c9.max_tags) |_| _ = try p.client.submit(.{ .stat = .{ .fid = 0 } }); + try testing.expectError(error.NoTags, p.client.submit(.{ .stat = .{ .fid = 0 } })); + const flush = try p.client.submit(.{ .flush = .{ .oldtag = 0 } }); + try testing.expectEqual(@as(u16, c9.max_tags), flush); +} + +test "multiple flushes retain reservations until each response and can themselves be flushed" { + var p: Pair = .{}; + try p.init(); + const oldtag = try p.client.submit(.{ .read = .{ .fid = 1, .offset = 0, .count = 1 } }); + _ = try p.nextRequest(); + p.server.release(); + const first = try p.client.submit(.{ .flush = .{ .oldtag = oldtag } }); + _ = try p.nextRequest(); + p.server.release(); + const second = try p.client.submit(.{ .flush = .{ .oldtag = oldtag } }); + _ = try p.nextRequest(); + p.server.release(); + try p.server.reply(first, .rflush); + _ = try p.nextResult(); + const held = try p.client.submit(.{ .stat = .{ .fid = 1 } }); + try testing.expect(held != oldtag); + try p.server.reply(second, .rflush); + _ = try p.nextResult(); + try testing.expectEqual(oldtag, try p.client.submit(.{ .stat = .{ .fid = 1 } })); + + try p.init(); + const read_tag = try p.client.submit(.{ .read = .{ .fid = 1, .offset = 0, .count = 1 } }); + _ = try p.nextRequest(); + p.server.release(); + const flush_tag = try p.client.submit(.{ .flush = .{ .oldtag = read_tag } }); + _ = try p.nextRequest(); + p.server.release(); + const cancel_flush = try p.client.submit(.{ .flush = .{ .oldtag = flush_tag } }); + _ = try p.nextRequest(); + p.server.release(); + try p.server.reply(read_tag, .{ .rread = .{ .data = "x" } }); + _ = try p.nextResult(); + try p.server.reply(cancel_flush, .rflush); + _ = try p.nextResult(); + try testing.expectEqual(@as(usize, 0), p.client.pending()); +} + +test "server reserves cancellation capacity and rejects duplicate pending tags" { + var input: [4096]u8 = undefined; + var output: [4096]u8 = undefined; + var server: c9.Server = .init(.{ .in = &input, .out = &output }); + server.msize = 4096; + var frame: [64]u8 = undefined; + for (0..64) |i| { + const bytes = try c9.encode(.{ .tstat = .{ .fid = 1 } }, @intCast(i), &frame); + _ = server.push(bytes); + _ = (try server.receive()).?; + server.release(); + } + _ = server.push(try c9.encode(.{ .tflush = .{ .oldtag = 0 } }, 100, &frame)); + _ = (try server.receive()).?; + server.release(); + try server.reply(100, .rflush); + server.wrote(server.output().len); + _ = server.push(try c9.encode(.{ .tstat = .{ .fid = 1 } }, 1, &frame)); + try testing.expectError(error.Protocol, server.receive()); + try testing.expect(server.dead); +} diff --git a/src/transport.zig b/src/transport.zig new file mode 100644 index 0000000..d7619be --- /dev/null +++ b/src/transport.zig @@ -0,0 +1,207 @@ +//! Stream transport adapters. Namespace paths, mounting, connection limits, and +//! event-loop scheduling belong to the application. POSIX descriptors are owned +//! by the caller; std.Io streams retain the standard library ownership contract. +const std = @import("std"); +const libc = std.c; +const wire = @import("wire.zig"); +pub const darwin = @import("builtin").os.tag.isDarwin(); +pub const sun_path_len = @typeInfo(@FieldType(libc.sockaddr.un, "path")).array.len; +pub const Address = union(enum) { unix: [:0]const u8, tcp: std.Io.net.IpAddress }; +pub const Error = error{ Socket, Flags, SocketOption, BadAddress, Bind, Listen, Connect, Closed, Timeout, Io }; + +/// Read one frame from any std.Io reader, including TCP, Unix sockets, or files. +pub fn readFrame(reader: *std.Io.Reader, buffer: []u8, msize: u32) ![]const u8 { + if (buffer.len < wire.header_len) return error.NoSpace; + try reader.readSliceAll(buffer[0..4]); + const len = wire.frameLen(buffer[0..4]).?; + if (len < wire.header_len) return error.BadValue; + if (len > msize or len > buffer.len) return error.Overlong; + try reader.readSliceAll(buffer[4..len]); + return buffer[0..len]; +} + +/// Batching and flushing are explicit: this does not flush the writer. +pub fn writeFrame(writer: *std.Io.Writer, frame: []const u8, msize: u32) !void { + if (frame.len > msize) return error.Overlong; + _ = try wire.decode(frame); + try writer.writeAll(frame); +} + +pub fn connect(io: std.Io, address: Address) !std.Io.net.Stream { + return switch (address) { + .tcp => |ip| ip.connect(io, .{ .mode = .stream }), + .unix => |path| (try std.Io.net.UnixAddress.init(path)).connect(io), + }; +} + +pub fn listen(io: std.Io, address: Address, backlog: u31) !std.Io.net.Server { + return switch (address) { + .tcp => |ip| ip.listen(io, .{ .kernel_backlog = backlog }), + .unix => |path| (try std.Io.net.UnixAddress.init(path)).listen(io, .{ .kernel_backlog = backlog }), + }; +} + +pub fn nowMs() i64 { + var ts: libc.timespec = undefined; + if (libc.clock_gettime(.MONOTONIC, &ts) != 0) return std.math.maxInt(i64); + return @as(i64, ts.sec) * std.time.ms_per_s + @divTrunc(ts.nsec, std.time.ns_per_ms); +} + +pub fn configure(fd: c_int, tcp: bool) Error!void { + const flags = libc.fcntl(fd, libc.F.GETFL, @as(c_int, 0)); + if (flags < 0) return error.Flags; + var options: libc.O = @bitCast(@as(u32, @bitCast(flags))); + options.NONBLOCK = true; + if (libc.fcntl(fd, libc.F.SETFL, @as(c_int, @bitCast(@as(u32, @bitCast(options))))) != 0 or + libc.fcntl(fd, libc.F.SETFD, @as(c_int, libc.FD_CLOEXEC)) != 0) return error.Flags; + const on: c_int = 1; + if (tcp and libc.setsockopt(fd, libc.IPPROTO.TCP, libc.TCP.NODELAY, &on, @sizeOf(c_int)) != 0) + return error.SocketOption; + if (comptime darwin) { + if (libc.setsockopt(fd, libc.SOL.SOCKET, libc.SO.NOSIGPIPE, &on, @sizeOf(c_int)) != 0) + return error.SocketOption; + } +} + +pub fn ipSockaddr(ip: std.Io.net.IpAddress, out: *libc.sockaddr.storage) libc.socklen_t { + return switch (ip) { + .ip4 => |a| blk: { + const addr: *libc.sockaddr.in = @ptrCast(out); + addr.* = .{ .port = std.mem.nativeToBig(u16, a.port), .addr = @bitCast(a.bytes) }; + break :blk @sizeOf(libc.sockaddr.in); + }, + .ip6 => |a| blk: { + const addr: *libc.sockaddr.in6 = @ptrCast(out); + addr.* = .{ .port = std.mem.nativeToBig(u16, a.port), .addr = a.bytes, .flowinfo = 0, .scope_id = a.interface.index }; + break :blk @sizeOf(libc.sockaddr.in6); + }, + }; +} + +fn sockaddr(address: Address, out: *libc.sockaddr.storage) Error!libc.socklen_t { + return switch (address) { + .tcp => |ip| ipSockaddr(ip, out), + .unix => |path| blk: { + if (path.len >= sun_path_len or std.mem.indexOfScalar(u8, path, 0) != null) + return error.BadAddress; + const addr: *libc.sockaddr.un = @ptrCast(out); + addr.* = .{ .path = @splat(0) }; + @memcpy(addr.path[0..path.len], path); + break :blk @sizeOf(libc.sockaddr.un); + }, + }; +} + +pub fn wait(fd: c_int, events: i16, deadline_ms: i64) Error!void { + while (true) { + const left = deadline_ms -| nowMs(); + if (left <= 0) return error.Timeout; + var fds = [1]libc.pollfd{.{ .fd = fd, .events = events, .revents = 0 }}; + const rc = libc.poll(&fds, 1, @intCast(@min(left, std.math.maxInt(c_int)))); + if (rc < 0) { + if (libc.errno(rc) == .INTR) continue; + return error.Io; + } + if (rc == 0) continue; + if (nowMs() >= deadline_ms) return error.Timeout; + if (fds[0].revents & events != 0) return; + return error.Closed; + } +} + +/// Connect a nonblocking socket with an absolute monotonic deadline. +pub fn connectFd(address: Address, deadline_ms: i64) Error!c_int { + var addr: libc.sockaddr.storage = undefined; + const len = try sockaddr(address, &addr); + const fd = libc.socket(addr.family, libc.SOCK.STREAM, 0); + if (fd < 0) return error.Socket; + errdefer close(fd); + try configure(fd, address == .tcp); + const rc = libc.connect(fd, @ptrCast(&addr), len); + if (rc != 0) { + switch (libc.errno(rc)) { + .INPROGRESS, .ALREADY, .INTR => {}, + else => return error.Connect, + } + try wait(fd, @intCast(libc.POLL.OUT), deadline_ms); + var status: c_int = 0; + var size: libc.socklen_t = @sizeOf(c_int); + if (libc.getsockopt(fd, libc.SOL.SOCKET, libc.SO.ERROR, @ptrCast(&status), &size) != 0) + return error.Connect; + if (status != 0) return error.Connect; + } + return fd; +} + +/// Does not unlink Unix paths or alter their permissions. +pub fn listenFd(address: Address, backlog: u31) Error!c_int { + var addr: libc.sockaddr.storage = undefined; + const len = try sockaddr(address, &addr); + const fd = libc.socket(addr.family, libc.SOCK.STREAM, 0); + if (fd < 0) return error.Socket; + errdefer close(fd); + try configure(fd, false); + if (address == .tcp) { + const on: c_int = 1; + if (libc.setsockopt(fd, libc.SOL.SOCKET, libc.SO.REUSEADDR, &on, @sizeOf(c_int)) != 0) + return error.SocketOption; + if (address.tcp == .ip6) { + const v6only = if (darwin) 27 else std.os.linux.IPV6.V6ONLY; + if (libc.setsockopt(fd, libc.IPPROTO.IPV6, v6only, &on, @sizeOf(c_int)) != 0) + return error.SocketOption; + } + } + if (libc.bind(fd, @ptrCast(&addr), len) != 0) return error.Bind; + if (libc.listen(fd, backlog) != 0) return error.Listen; + return fd; +} + +pub fn acceptFd(listener: c_int, tcp: bool) Error!?c_int { + const fd = libc.accept(listener, null, null); + if (fd < 0) return switch (libc.errno(fd)) { + .AGAIN, .INTR, .CONNABORTED => null, + else => error.Socket, + }; + errdefer close(fd); + try configure(fd, tcp); + return fd; +} + +/// null means retry after readiness; zero means EOF. Empty buffers are forbidden. +pub fn read(fd: c_int, buffer: []u8) Error!?usize { + std.debug.assert(buffer.len > 0); + const n = libc.read(fd, buffer.ptr, buffer.len); + if (n < 0) return switch (libc.errno(n)) { + .INTR, .AGAIN => null, + else => error.Io, + }; + return @intCast(n); +} + +pub fn write(fd: c_int, bytes: []const u8) Error!?usize { + std.debug.assert(bytes.len > 0); + const n = libc.send(fd, bytes.ptr, bytes.len, if (darwin) 0 else libc.MSG.NOSIGNAL); + if (n < 0) return switch (libc.errno(n)) { + .INTR, .AGAIN => null, + else => error.Io, + }; + if (n == 0) return error.Closed; + return @intCast(n); +} + +pub fn close(fd: c_int) void { + _ = libc.close(fd); +} + +/// Probe a Unix listener without changing namespace entries. Uncertainty is live. +pub fn isListening(path: [:0]const u8) bool { + var addr: libc.sockaddr.un = .{ .path = @splat(0) }; + if (path.len + 1 > sun_path_len) return true; // cannot ask; assume occupied + @memcpy(addr.path[0 .. path.len + 1], path[0 .. path.len + 1]); + const fd = libc.socket(libc.AF.UNIX, libc.SOCK.STREAM, 0); + if (fd < 0) return true; + defer _ = libc.close(fd); + configure(fd, false) catch return true; + if (libc.connect(fd, @ptrCast(&addr), @sizeOf(@TypeOf(addr))) == 0) return true; + return libc.errno(-1) != .CONNREFUSED; +} diff --git a/src/wire.zig b/src/wire.zig new file mode 100644 index 0000000..e093450 --- /dev/null +++ b/src/wire.zig @@ -0,0 +1,1153 @@ +//! Base 9P2000 wire format. Decoded strings and data borrow the input frame. +//! Encoding performs a complete size check before writing the caller-owned buffer. +const std = @import("std"); +const assert = std.debug.assert; + +pub const Error = error{ + Truncated, + Overlong, + BadTag, + BadValue, + Trailing, + NoSpace, +}; + +pub const Type = enum(u8) { + tversion = 100, + rversion = 101, + tauth = 102, + rauth = 103, + tattach = 104, + rattach = 105, + terror = 106, + rerror = 107, + tflush = 108, + rflush = 109, + twalk = 110, + rwalk = 111, + topen = 112, + ropen = 113, + tcreate = 114, + rcreate = 115, + tread = 116, + rread = 117, + twrite = 118, + rwrite = 119, + tclunk = 120, + rclunk = 121, + tremove = 122, + rremove = 123, + tstat = 124, + rstat = 125, + twstat = 126, + rwstat = 127, + _, +}; + +pub fn isT(t: Type) bool { + return @intFromEnum(t) % 2 == 0; +} + +pub const header_len: usize = 4 + 1 + 2; + +pub const qid_len: usize = 1 + 4 + 8; + +pub const stat_fixed: usize = 2 + qid_len + 5 * 2 + 4 * 4 + 8; + +pub const notag: u16 = 0xFFFF; + +pub const nofid: u32 = 0xFFFF_FFFF; + +pub const max_welem: usize = 16; + +const test_msize: u32 = 4096; + +pub const iohdrsz: u32 = 24; + +pub const qtdir: u8 = 0x80; +pub const qtappend: u8 = 0x40; +pub const qtexcl: u8 = 0x20; +pub const qtmount: u8 = 0x10; +pub const qtauth: u8 = 0x08; +pub const qttmp: u8 = 0x04; +pub const qtfile: u8 = 0x00; + +pub const dmdir: u32 = 0x8000_0000; +pub const dmappend: u32 = 0x4000_0000; +pub const dmexcl: u32 = 0x2000_0000; +pub const dmmount: u32 = 0x1000_0000; +pub const dmauth: u32 = 0x0800_0000; +pub const dmtmp: u32 = 0x0400_0000; +pub const dmperm: u32 = 0o777; + +comptime { + assert(header_len == 7); + assert(qid_len == 13); + assert(stat_fixed == 49); + for (std.enums.values(Type)) |t| { + const even = @intFromEnum(t) % 2 == 0; + assert(isT(t) == even); + assert(std.mem.startsWith(u8, @tagName(t), if (even) "t" else "r")); + } +} + +pub const Qid = struct { + type: u8, + version: u32, + path: u64, + + pub fn encode(self: Qid, buf: []u8) Error![]u8 { + if (buf.len < qid_len) return error.NoSpace; + buf[0] = self.type; + std.mem.writeInt(u32, buf[1..5], self.version, .little); + std.mem.writeInt(u64, buf[5..13], self.path, .little); + return buf[0..qid_len]; + } + + pub fn decode(bytes: []const u8) Error!Qid { + if (bytes.len < qid_len) return error.Truncated; + return .{ + .type = bytes[0], + .version = std.mem.readInt(u32, bytes[1..5], .little), + .path = std.mem.readInt(u64, bytes[5..13], .little), + }; + } +}; + +pub const Stat = struct { + type: u16, + dev: u32, + qid: Qid, + mode: u32, + atime: u32, + mtime: u32, + length: u64, + name: []const u8, + uid: []const u8, + gid: []const u8, + muid: []const u8, + + pub fn size(self: Stat) Error!u16 { + const n = + 2 + // type + 4 + // dev + qid_len + // qid: type[1] version[4] path[8] + 4 + // mode + 4 + // atime + 4 + // mtime + 8 + // length + try stringLen(self.name) + + try stringLen(self.uid) + + try stringLen(self.gid) + + try stringLen(self.muid); + assert(n >= stat_fixed - 2); + if (n > std.math.maxInt(u16)) return error.Overlong; + return @intCast(n); + } + + pub fn encode(self: Stat, buf: []u8) Error![]u8 { + const n = try self.size(); + const total = @as(usize, n) + 2; + if (buf.len < total) return error.NoSpace; + var w: Writer = .init(buf[0..total]); + try w.putU16(n); + try w.putU16(self.type); + try w.putU32(self.dev); + try w.putQid(self.qid); + try w.putU32(self.mode); + try w.putU32(self.atime); + try w.putU32(self.mtime); + try w.putU64(self.length); + try w.putString(self.name); + try w.putString(self.uid); + try w.putString(self.gid); + try w.putString(self.muid); + assert(w.n == total); + return buf[0..total]; + } + + pub fn decode(bytes: []const u8) Error!Stat { + var r: Reader = .init(bytes); + const n = try r.getU16(); + const body = bytes.len - 2; + if (n > body) return error.Truncated; + if (n < body) return error.Trailing; + const self: Stat = .{ + .type = try r.getU16(), + .dev = try r.getU32(), + .qid = try r.getQid(), + .mode = try r.getU32(), + .atime = try r.getU32(), + .mtime = try r.getU32(), + .length = try r.getU64(), + .name = try r.getString(), + .uid = try r.getString(), + .gid = try r.getString(), + .muid = try r.getString(), + }; + try r.end(); + return self; + } +}; + +pub const Msg = union(enum) { + tversion: struct { msize: u32, version: []const u8 }, + rversion: struct { msize: u32, version: []const u8 }, + + tauth: struct { afid: u32, uname: []const u8, aname: []const u8 }, + rauth: struct { aqid: Qid }, + + tattach: struct { fid: u32, afid: u32, uname: []const u8, aname: []const u8 }, + rattach: struct { qid: Qid }, + + rerror: struct { ename: []const u8 }, + + tflush: struct { oldtag: u16 }, + rflush: void, + + twalk: struct { + fid: u32, + newfid: u32, + nwname: u16, + wname: [max_welem][]const u8 = @splat(""), + }, + rwalk: struct { + nwqid: u16, + wqid: [max_welem]Qid = @splat(.{ .type = 0, .version = 0, .path = 0 }), + }, + + topen: struct { fid: u32, mode: u8 }, + ropen: struct { qid: Qid, iounit: u32 }, + + tcreate: struct { fid: u32, name: []const u8, perm: u32, mode: u8 }, + rcreate: struct { qid: Qid, iounit: u32 }, + + tread: struct { fid: u32, offset: u64, count: u32 }, + rread: struct { data: []const u8 }, + + twrite: struct { fid: u32, offset: u64, data: []const u8 }, + rwrite: struct { count: u32 }, + + tclunk: struct { fid: u32 }, + rclunk: void, + tremove: struct { fid: u32 }, + rremove: void, + + tstat: struct { fid: u32 }, + rstat: struct { stat: Stat }, + twstat: struct { fid: u32, stat: Stat }, + rwstat: void, + + pub fn msgType(msg: Msg) Type { + return switch (msg) { + inline else => |_, t| @field(Type, @tagName(t)), + }; + } +}; + +pub const Decoded = struct { + tag: u16, + msg: Msg, +}; + +pub fn frameLen(prefix: []const u8) ?u32 { + if (prefix.len < 4) return null; + return std.mem.readInt(u32, prefix[0..4], .little); +} + +pub fn encodedLen(msg: Msg) Error!usize { + const body: u64 = switch (msg) { + .tversion => |m| 4 + try stringLen(m.version), + .rversion => |m| 4 + try stringLen(m.version), + .tauth => |m| 4 + try stringLen(m.uname) + try stringLen(m.aname), + .rauth => qid_len, + .tattach => |m| 4 + 4 + try stringLen(m.uname) + try stringLen(m.aname), + .rattach => qid_len, + .rerror => |m| try stringLen(m.ename), + .tflush => 2, + .rflush => 0, + .twalk => |m| blk: { + if (m.nwname > max_welem) return error.Overlong; + var n: usize = 4 + 4 + 2; + for (m.wname[0..m.nwname]) |name| n += try stringLen(name); + break :blk n; + }, + .rwalk => |m| blk: { + if (m.nwqid > max_welem) return error.Overlong; + break :blk 2 + @as(usize, m.nwqid) * qid_len; + }, + .topen => 4 + 1, + .ropen => qid_len + 4, + .tcreate => |m| 4 + try stringLen(m.name) + 4 + 1, + .rcreate => qid_len + 4, + .tread => 4 + 8 + 4, + .rread => |m| try dataLen(m.data), + .twrite => |m| 4 + 8 + try dataLen(m.data), + .rwrite => 4, + .tclunk => 4, + .rclunk => 0, + .tremove => 4, + .rremove => 0, + .tstat => 4, + .rstat => |m| 2 + try statLen(m.stat), + .twstat => |m| 4 + 2 + try statLen(m.stat), + .rwstat => 0, + }; + const total = header_len + body; + if (total > std.math.maxInt(u32)) return error.Overlong; + return @intCast(total); +} + +fn statLen(stat: Stat) Error!usize { + const n = @as(usize, try stat.size()) + 2; + if (n > std.math.maxInt(u16)) return error.Overlong; + return n; +} + +fn stringLen(s: []const u8) Error!usize { + if (s.len > std.math.maxInt(u16)) return error.Overlong; + return 2 + s.len; +} + +fn dataLen(d: []const u8) Error!u64 { + if (d.len > std.math.maxInt(u32)) return error.Overlong; + return 4 + @as(u64, d.len); +} + +pub fn encode(msg: Msg, tag: u16, buf: []u8) Error![]u8 { + const total = try encodedLen(msg); + if (total > buf.len) return error.NoSpace; + + var w: Writer = .init(buf[0..total]); + try w.putU32(@intCast(total)); + try w.putByte(@intFromEnum(msg.msgType())); + try w.putU16(tag); + + switch (msg) { + .tversion => |m| { + try w.putU32(m.msize); + try w.putString(m.version); + }, + .rversion => |m| { + try w.putU32(m.msize); + try w.putString(m.version); + }, + .tauth => |m| { + try w.putU32(m.afid); + try w.putString(m.uname); + try w.putString(m.aname); + }, + .rauth => |m| try w.putQid(m.aqid), + .tattach => |m| { + try w.putU32(m.fid); + try w.putU32(m.afid); + try w.putString(m.uname); + try w.putString(m.aname); + }, + .rattach => |m| try w.putQid(m.qid), + .rerror => |m| try w.putString(m.ename), + .tflush => |m| try w.putU16(m.oldtag), + .rflush => {}, + .twalk => |m| { + try w.putU32(m.fid); + try w.putU32(m.newfid); + try w.putU16(m.nwname); + for (m.wname[0..m.nwname]) |name| try w.putString(name); + }, + .rwalk => |m| { + try w.putU16(m.nwqid); + for (m.wqid[0..m.nwqid]) |qid| try w.putQid(qid); + }, + .topen => |m| { + try w.putU32(m.fid); + try w.putByte(m.mode); + }, + .ropen => |m| { + try w.putQid(m.qid); + try w.putU32(m.iounit); + }, + .tcreate => |m| { + try w.putU32(m.fid); + try w.putString(m.name); + try w.putU32(m.perm); + try w.putByte(m.mode); + }, + .rcreate => |m| { + try w.putQid(m.qid); + try w.putU32(m.iounit); + }, + .tread => |m| { + try w.putU32(m.fid); + try w.putU64(m.offset); + try w.putU32(m.count); + }, + .rread => |m| { + try w.putU32(@intCast(m.data.len)); + try w.putBytes(m.data); + }, + .twrite => |m| { + try w.putU32(m.fid); + try w.putU64(m.offset); + try w.putU32(@intCast(m.data.len)); + try w.putBytes(m.data); + }, + .rwrite => |m| try w.putU32(m.count), + .tclunk => |m| try w.putU32(m.fid), + .rclunk => {}, + .tremove => |m| try w.putU32(m.fid), + .rremove => {}, + .tstat => |m| try w.putU32(m.fid), + .rstat => |m| { + try w.putU16(try m.stat.size() + 2); + try w.putStat(m.stat); + }, + .twstat => |m| { + try w.putU32(m.fid); + try w.putU16(try m.stat.size() + 2); + try w.putStat(m.stat); + }, + .rwstat => {}, + } + + assert(w.n == total); + return buf[0..total]; +} + +pub fn decode(bytes: []const u8) Error!Decoded { + if (bytes.len < header_len) return error.Truncated; + const size = std.mem.readInt(u32, bytes[0..4], .little); + if (size < header_len) return error.BadValue; + if (size > bytes.len) return error.Truncated; + if (size < bytes.len) return error.Trailing; + + const t: Type = @enumFromInt(bytes[4]); + const tag = std.mem.readInt(u16, bytes[5..7], .little); + + var r: Reader = .init(bytes[header_len..size]); + + const msg: Msg = switch (t) { + .tversion => .{ .tversion = .{ .msize = try r.getU32(), .version = try r.getString() } }, + .rversion => .{ .rversion = .{ .msize = try r.getU32(), .version = try r.getString() } }, + .tauth => .{ .tauth = .{ + .afid = try r.getU32(), + .uname = try r.getString(), + .aname = try r.getString(), + } }, + .rauth => .{ .rauth = .{ .aqid = try r.getQid() } }, + .tattach => .{ .tattach = .{ + .fid = try r.getU32(), + .afid = try r.getU32(), + .uname = try r.getString(), + .aname = try r.getString(), + } }, + .rattach => .{ .rattach = .{ .qid = try r.getQid() } }, + .rerror => .{ .rerror = .{ .ename = try r.getString() } }, + .tflush => .{ .tflush = .{ .oldtag = try r.getU16() } }, + .rflush => .rflush, + .twalk => blk: { + var m: Msg = .{ .twalk = .{ + .fid = try r.getU32(), + .newfid = try r.getU32(), + .nwname = try r.getU16(), + } }; + if (m.twalk.nwname > max_welem) return error.Overlong; + for (m.twalk.wname[0..m.twalk.nwname]) |*name| name.* = try r.getString(); + break :blk m; + }, + .rwalk => blk: { + var m: Msg = .{ .rwalk = .{ .nwqid = try r.getU16() } }; + if (m.rwalk.nwqid > max_welem) return error.Overlong; + for (m.rwalk.wqid[0..m.rwalk.nwqid]) |*qid| qid.* = try r.getQid(); + break :blk m; + }, + .topen => .{ .topen = .{ .fid = try r.getU32(), .mode = try r.getByte() } }, + .ropen => .{ .ropen = .{ .qid = try r.getQid(), .iounit = try r.getU32() } }, + .tcreate => .{ .tcreate = .{ + .fid = try r.getU32(), + .name = try r.getString(), + .perm = try r.getU32(), + .mode = try r.getByte(), + } }, + .rcreate => .{ .rcreate = .{ .qid = try r.getQid(), .iounit = try r.getU32() } }, + .tread => .{ .tread = .{ + .fid = try r.getU32(), + .offset = try r.getU64(), + .count = try r.getU32(), + } }, + .rread => .{ .rread = .{ .data = try r.getData() } }, + .twrite => .{ .twrite = .{ + .fid = try r.getU32(), + .offset = try r.getU64(), + .data = try r.getData(), + } }, + .rwrite => .{ .rwrite = .{ .count = try r.getU32() } }, + .tclunk => .{ .tclunk = .{ .fid = try r.getU32() } }, + .rclunk => .rclunk, + .tremove => .{ .tremove = .{ .fid = try r.getU32() } }, + .rremove => .rremove, + .tstat => .{ .tstat = .{ .fid = try r.getU32() } }, + .rstat => .{ .rstat = .{ .stat = try Stat.decode(try r.getBlob16()) } }, + .twstat => .{ .twstat = .{ + .fid = try r.getU32(), + .stat = try Stat.decode(try r.getBlob16()), + } }, + .rwstat => .rwstat, + .terror, _ => return error.BadTag, + }; + + try r.end(); + return .{ .tag = tag, .msg = msg }; +} + +const Writer = struct { + buf: []u8, + n: usize = 0, + + fn init(buf: []u8) Writer { + return .{ .buf = buf }; + } + + fn room(w: *Writer, k: usize) Error![]u8 { + if (w.buf.len - w.n < k) return error.NoSpace; + defer w.n += k; + return w.buf[w.n..][0..k]; + } + + fn putByte(w: *Writer, v: u8) Error!void { + (try w.room(1))[0] = v; + } + + fn putU16(w: *Writer, v: u16) Error!void { + std.mem.writeInt(u16, (try w.room(2))[0..2], v, .little); + } + + fn putU32(w: *Writer, v: u32) Error!void { + std.mem.writeInt(u32, (try w.room(4))[0..4], v, .little); + } + + fn putU64(w: *Writer, v: u64) Error!void { + std.mem.writeInt(u64, (try w.room(8))[0..8], v, .little); + } + + fn putBytes(w: *Writer, v: []const u8) Error!void { + const target = try w.room(v.len); + // Permit payloads staged at their final position in the output frame. + if (target.ptr != v.ptr) @memcpy(target, v); + } + + fn putString(w: *Writer, v: []const u8) Error!void { + assert(v.len <= std.math.maxInt(u16)); + try w.putU16(@intCast(v.len)); + try w.putBytes(v); + } + + fn putQid(w: *Writer, v: Qid) Error!void { + _ = try v.encode(try w.room(qid_len)); + } + + fn putStat(w: *Writer, v: Stat) Error!void { + const total = @as(usize, try v.size()) + 2; + _ = try v.encode(try w.room(total)); + } +}; + +const Reader = struct { + bytes: []const u8, + i: usize = 0, + + fn init(bytes: []const u8) Reader { + return .{ .bytes = bytes }; + } + + fn take(r: *Reader, n: usize) Error![]const u8 { + if (r.bytes.len - r.i < n) return error.Truncated; + defer r.i += n; + return r.bytes[r.i..][0..n]; + } + + fn getByte(r: *Reader) Error!u8 { + return (try r.take(1))[0]; + } + + fn getU16(r: *Reader) Error!u16 { + return std.mem.readInt(u16, (try r.take(2))[0..2], .little); + } + + fn getU32(r: *Reader) Error!u32 { + return std.mem.readInt(u32, (try r.take(4))[0..4], .little); + } + + fn getU64(r: *Reader) Error!u64 { + return std.mem.readInt(u64, (try r.take(8))[0..8], .little); + } + + fn getString(r: *Reader) Error![]const u8 { + return r.take(try r.getU16()); + } + + fn getData(r: *Reader) Error![]const u8 { + return r.take(try r.getU32()); + } + + fn getBlob16(r: *Reader) Error![]const u8 { + return r.take(try r.getU16()); + } + + fn getQid(r: *Reader) Error!Qid { + return Qid.decode(try r.take(qid_len)); + } + + fn end(r: *Reader) Error!void { + if (r.i != r.bytes.len) return error.Trailing; + } +}; + +const testing = std.testing; + +fn roundTrip(buf: []u8, tag: u16, msg: Msg) !Msg { + const bytes = try encode(msg, tag, buf); + try testing.expectEqual(bytes.len, frameLen(bytes).?); + const got = try decode(bytes); + try testing.expectEqual(tag, got.tag); + try testing.expectEqual(msg.msgType(), got.msg.msgType()); + try expectMsgEqual(msg, got.msg); + return got.msg; +} + +fn expectStatEqual(want: Stat, have: Stat) !void { + try testing.expectEqual(want.type, have.type); + try testing.expectEqual(want.dev, have.dev); + try testing.expectEqual(want.qid, have.qid); + try testing.expectEqual(want.mode, have.mode); + try testing.expectEqual(want.atime, have.atime); + try testing.expectEqual(want.mtime, have.mtime); + try testing.expectEqual(want.length, have.length); + try testing.expectEqualStrings(want.name, have.name); + try testing.expectEqualStrings(want.uid, have.uid); + try testing.expectEqualStrings(want.gid, have.gid); + try testing.expectEqualStrings(want.muid, have.muid); +} + +fn expectMsgEqual(want: Msg, have: Msg) !void { + switch (want) { + .tversion => |w| { + try testing.expectEqual(w.msize, have.tversion.msize); + try testing.expectEqualStrings(w.version, have.tversion.version); + }, + .rversion => |w| { + try testing.expectEqual(w.msize, have.rversion.msize); + try testing.expectEqualStrings(w.version, have.rversion.version); + }, + .tauth => |w| { + try testing.expectEqual(w.afid, have.tauth.afid); + try testing.expectEqualStrings(w.uname, have.tauth.uname); + try testing.expectEqualStrings(w.aname, have.tauth.aname); + }, + .rauth => |w| try testing.expectEqual(w.aqid, have.rauth.aqid), + .tattach => |w| { + try testing.expectEqual(w.fid, have.tattach.fid); + try testing.expectEqual(w.afid, have.tattach.afid); + try testing.expectEqualStrings(w.uname, have.tattach.uname); + try testing.expectEqualStrings(w.aname, have.tattach.aname); + }, + .rattach => |w| try testing.expectEqual(w.qid, have.rattach.qid), + .rerror => |w| try testing.expectEqualStrings(w.ename, have.rerror.ename), + .tflush => |w| try testing.expectEqual(w.oldtag, have.tflush.oldtag), + .rflush, .rclunk, .rremove, .rwstat => {}, + .twalk => |w| { + try testing.expectEqual(w.fid, have.twalk.fid); + try testing.expectEqual(w.newfid, have.twalk.newfid); + try testing.expectEqual(w.nwname, have.twalk.nwname); + for (w.wname[0..w.nwname], have.twalk.wname[0..w.nwname]) |a, b| + try testing.expectEqualStrings(a, b); + }, + .rwalk => |w| { + try testing.expectEqual(w.nwqid, have.rwalk.nwqid); + for (w.wqid[0..w.nwqid], have.rwalk.wqid[0..w.nwqid]) |a, b| + try testing.expectEqual(a, b); + }, + .topen => |w| { + try testing.expectEqual(w.fid, have.topen.fid); + try testing.expectEqual(w.mode, have.topen.mode); + }, + .ropen => |w| { + try testing.expectEqual(w.qid, have.ropen.qid); + try testing.expectEqual(w.iounit, have.ropen.iounit); + }, + .tcreate => |w| { + try testing.expectEqual(w.fid, have.tcreate.fid); + try testing.expectEqualStrings(w.name, have.tcreate.name); + try testing.expectEqual(w.perm, have.tcreate.perm); + try testing.expectEqual(w.mode, have.tcreate.mode); + }, + .rcreate => |w| { + try testing.expectEqual(w.qid, have.rcreate.qid); + try testing.expectEqual(w.iounit, have.rcreate.iounit); + }, + .tread => |w| { + try testing.expectEqual(w.fid, have.tread.fid); + try testing.expectEqual(w.offset, have.tread.offset); + try testing.expectEqual(w.count, have.tread.count); + }, + .rread => |w| try testing.expectEqualStrings(w.data, have.rread.data), + .twrite => |w| { + try testing.expectEqual(w.fid, have.twrite.fid); + try testing.expectEqual(w.offset, have.twrite.offset); + try testing.expectEqualStrings(w.data, have.twrite.data); + }, + .rwrite => |w| try testing.expectEqual(w.count, have.rwrite.count), + .tclunk => |w| try testing.expectEqual(w.fid, have.tclunk.fid), + .tremove => |w| try testing.expectEqual(w.fid, have.tremove.fid), + .tstat => |w| try testing.expectEqual(w.fid, have.tstat.fid), + .rstat => |w| try expectStatEqual(w.stat, have.rstat.stat), + .twstat => |w| { + try testing.expectEqual(w.fid, have.twstat.fid); + try expectStatEqual(w.stat, have.twstat.stat); + }, + } +} + +const sample_qid: Qid = .{ .type = qtdir, .version = 3, .path = 0x0102_0304_0506_0708 }; + +const sample_stat: Stat = .{ + .type = 0, + .dev = 0, + .qid = sample_qid, + .mode = dmdir | 0o755, + .atime = 1, + .mtime = 2, + .length = 0, + .name = "body", + .uid = "goblin", + .gid = "goblin", + .muid = "goblin", +}; + +test "9p: the type numbers and their parity are the protocol's own" { + try testing.expectEqual(@as(u8, 100), @intFromEnum(Type.tversion)); + try testing.expectEqual(@as(u8, 106), @intFromEnum(Type.terror)); + try testing.expectEqual(@as(u8, 107), @intFromEnum(Type.rerror)); + try testing.expectEqual(@as(u8, 126), @intFromEnum(Type.twstat)); + try testing.expectEqual(@as(u8, 127), @intFromEnum(Type.rwstat)); + + try testing.expectEqual(@as(usize, 28), std.enums.values(Type).len); + for (std.enums.values(Type), 100..) |t, want| try testing.expectEqual(@as(u8, @intCast(want)), @intFromEnum(t)); + + try testing.expect(isT(.tversion)); + try testing.expect(!isT(.rversion)); + try testing.expect(isT(.twstat)); + try testing.expect(!isT(.rwstat)); + + try testing.expectEqual(@as(u16, 0xFFFF), notag); + try testing.expectEqual(@as(u32, 0xFFFF_FFFF), nofid); + try testing.expectEqual(@as(usize, 16), max_welem); +} + +test "9p: a qid is thirteen bytes" { + var buf: [32]u8 = undefined; + const bytes = try sample_qid.encode(&buf); + try testing.expectEqual(qid_len, bytes.len); + try testing.expectEqual(@as(usize, 13), bytes.len); + try testing.expectEqual(sample_qid, try Qid.decode(bytes)); + try testing.expectError(error.Truncated, Qid.decode(bytes[0..12])); + try testing.expectError(error.NoSpace, sample_qid.encode(buf[0..12])); +} + +test "9p: an encoded stat is size() + 2 bytes" { + var buf: [256]u8 = undefined; + const bytes = try sample_stat.encode(&buf); + const n = try sample_stat.size(); + try testing.expectEqual(@as(usize, n) + 2, bytes.len); + try testing.expectEqual(@as(u16, 69), n); + try testing.expectEqual(stat_fixed - 2 + 22, n); + try testing.expectEqual(n, std.mem.readInt(u16, bytes[0..2], .little)); + try expectStatEqual(sample_stat, try Stat.decode(bytes)); + + const bare: Stat = .{ + .type = 0, + .dev = 0, + .qid = .{ .type = qtfile, .version = 0, .path = 0 }, + .mode = 0, + .atime = 0, + .mtime = 0, + .length = 0, + .name = "", + .uid = "", + .gid = "", + .muid = "", + }; + try testing.expectEqual(@as(u16, 47), try bare.size()); + try testing.expectEqual(@as(usize, 49), (try bare.encode(&buf)).len); +} + +test "9p: every message round-trips" { + var buf: [512]u8 = undefined; + + _ = try roundTrip(&buf, notag, .{ .tversion = .{ .msize = 8192, .version = "9P2000" } }); + _ = try roundTrip(&buf, notag, .{ .rversion = .{ .msize = 8192, .version = "9P2000" } }); + _ = try roundTrip(&buf, notag, .{ .rversion = .{ .msize = test_msize, .version = "unknown" } }); + _ = try roundTrip(&buf, 1, .{ .tauth = .{ .afid = 1, .uname = "goblin", .aname = "" } }); + _ = try roundTrip(&buf, 1, .{ .rauth = .{ .aqid = .{ .type = qtauth, .version = 0, .path = 9 } } }); + _ = try roundTrip(&buf, 2, .{ .tattach = .{ .fid = 0, .afid = nofid, .uname = "goblin", .aname = "" } }); + _ = try roundTrip(&buf, 2, .{ .rattach = .{ .qid = sample_qid } }); + _ = try roundTrip(&buf, 3, .{ .rerror = .{ .ename = "no such file" } }); + _ = try roundTrip(&buf, 4, .{ .tflush = .{ .oldtag = 3 } }); + _ = try roundTrip(&buf, 4, .rflush); + _ = try roundTrip(&buf, 5, .{ .twalk = .{ .fid = 0, .newfid = 1, .nwname = 2, .wname = .{ "7", "body" } ++ @as([max_welem - 2][]const u8, @splat("")) } }); + _ = try roundTrip(&buf, 5, .{ .rwalk = .{ .nwqid = 2, .wqid = .{ sample_qid, sample_qid } ++ @as([max_welem - 2]Qid, @splat(sample_qid)) } }); + _ = try roundTrip(&buf, 6, .{ .topen = .{ .fid = 1, .mode = 0 } }); + _ = try roundTrip(&buf, 6, .{ .ropen = .{ .qid = sample_qid, .iounit = 8192 - iohdrsz } }); + _ = try roundTrip(&buf, 7, .{ .tcreate = .{ .fid = 1, .name = "new", .perm = dmdir | 0o777, .mode = 2 } }); + _ = try roundTrip(&buf, 7, .{ .rcreate = .{ .qid = sample_qid, .iounit = 0 } }); + _ = try roundTrip(&buf, 8, .{ .tread = .{ .fid = 1, .offset = 0xdead_beef_cafe, .count = 4096 } }); + _ = try roundTrip(&buf, 8, .{ .rread = .{ .data = "hello" } }); + _ = try roundTrip(&buf, 8, .{ .rread = .{ .data = "" } }); + _ = try roundTrip(&buf, 9, .{ .twrite = .{ .fid = 1, .offset = 0, .data = "Edit ,d" } }); + _ = try roundTrip(&buf, 9, .{ .twrite = .{ .fid = 1, .offset = 0, .data = "" } }); + _ = try roundTrip(&buf, 9, .{ .rwrite = .{ .count = 7 } }); + _ = try roundTrip(&buf, 10, .{ .tclunk = .{ .fid = 1 } }); + _ = try roundTrip(&buf, 10, .rclunk); + _ = try roundTrip(&buf, 11, .{ .tremove = .{ .fid = 1 } }); + _ = try roundTrip(&buf, 11, .rremove); + _ = try roundTrip(&buf, 12, .{ .tstat = .{ .fid = 1 } }); + _ = try roundTrip(&buf, 12, .{ .rstat = .{ .stat = sample_stat } }); + _ = try roundTrip(&buf, 13, .{ .twstat = .{ .fid = 1, .stat = sample_stat } }); + _ = try roundTrip(&buf, 13, .rwstat); + + try testing.expectEqual(@as(usize, 27), @typeInfo(Msg).@"union".fields.len); + try testing.expectEqual(std.enums.values(Type).len - 1, @typeInfo(Msg).@"union".fields.len); +} + +test "9p: empty and maximum-length strings survive the trip" { + var buf: [70_000]u8 = undefined; + + const empty = try roundTrip(&buf, 1, .{ .tattach = .{ .fid = 0, .afid = nofid, .uname = "", .aname = "" } }); + try testing.expectEqual(@as(usize, 0), empty.tattach.uname.len); + try testing.expectEqual(@as(usize, header_len + 4 + 4 + 2 + 2), (try encode(empty, 1, &buf)).len); + + var big: [65_536]u8 = undefined; + @memset(&big, 'x'); + const max = big[0..std.math.maxInt(u16)]; + const got = try roundTrip(&buf, 1, .{ .rerror = .{ .ename = max } }); + try testing.expectEqual(@as(usize, 65_535), got.rerror.ename.len); + try testing.expectError(error.Overlong, encode(.{ .rerror = .{ .ename = &big } }, 1, &buf)); + + var wide = sample_stat; + wide.name = max; + try testing.expectError(error.Overlong, wide.size()); + try testing.expectError(error.Overlong, encode(.{ .rstat = .{ .stat = wide } }, 1, &buf)); +} + +test "9p: Twalk carries 0, 1 and 16 elements and refuses 17" { + var buf: [512]u8 = undefined; + + const zero = try roundTrip(&buf, 1, .{ .twalk = .{ .fid = 0, .newfid = 1, .nwname = 0 } }); + try testing.expectEqual(@as(u16, 0), zero.twalk.nwname); + try testing.expectEqual(@as(usize, header_len + 4 + 4 + 2), (try encode(zero, 1, &buf)).len); + + _ = try roundTrip(&buf, 1, .{ .twalk = .{ + .fid = 0, + .newfid = 1, + .nwname = 1, + .wname = .{"body"} ++ @as([max_welem - 1][]const u8, @splat("")), + } }); + + const names: [max_welem][]const u8 = .{ "a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l", "m", "n", "o", "p" }; + const full = try roundTrip(&buf, 1, .{ .twalk = .{ .fid = 0, .newfid = 1, .nwname = max_welem, .wname = names } }); + try testing.expectEqual(@as(u16, 16), full.twalk.nwname); + for (names, full.twalk.wname[0..max_welem]) |a, b| try testing.expectEqualStrings(a, b); + _ = try roundTrip(&buf, 1, .{ .rwalk = .{ .nwqid = max_welem, .wqid = @splat(sample_qid) } }); + + var raw: [256]u8 = undefined; + const bad = blk: { + var w: Writer = .init(&raw); + try w.putU32(0); // patched below + try w.putByte(@intFromEnum(Type.twalk)); + try w.putU16(1); + try w.putU32(0); + try w.putU32(1); + try w.putU16(17); + for (0..17) |i| try w.putString(&[_]u8{@intCast('a' + i)}); + std.mem.writeInt(u32, raw[0..4], @intCast(w.n), .little); + break :blk raw[0..w.n]; + }; + try testing.expectEqual(@as(usize, header_len + 4 + 4 + 2 + 17 * 3), bad.len); + try testing.expectError(error.Overlong, decode(bad)); + + const bad_r = blk: { + var w: Writer = .init(&raw); + try w.putU32(0); + try w.putByte(@intFromEnum(Type.rwalk)); + try w.putU16(1); + try w.putU16(17); + for (0..17) |_| try w.putQid(sample_qid); + std.mem.writeInt(u32, raw[0..4], @intCast(w.n), .little); + break :blk raw[0..w.n]; + }; + try testing.expectError(error.Overlong, decode(bad_r)); +} + +test "9p: the stat double length" { + var buf: [512]u8 = undefined; + + var good: [512]u8 = undefined; + const n = blk: { + const bytes = try encode(.{ .rstat = .{ .stat = sample_stat } }, 1, &buf); + @memcpy(good[0..bytes.len], bytes); + break :blk bytes.len; + }; + const inner = try sample_stat.size(); + try testing.expectEqual(inner + 2, std.mem.readInt(u16, good[header_len..][0..2], .little)); + try testing.expectEqual(inner, std.mem.readInt(u16, good[header_len + 2 ..][0..2], .little)); + try testing.expectEqual(header_len + 2 + @as(usize, inner) + 2, n); + + const w_bytes = try encode(.{ .twstat = .{ .fid = 7, .stat = sample_stat } }, 1, &buf); + try testing.expectEqual(inner + 2, std.mem.readInt(u16, w_bytes[header_len + 4 ..][0..2], .little)); + try testing.expectEqual(inner, std.mem.readInt(u16, w_bytes[header_len + 6 ..][0..2], .little)); + + var off: [512]u8 = undefined; + + @memcpy(off[0..n], good[0..n]); + std.mem.writeInt(u16, off[header_len..][0..2], inner, .little); + try testing.expectError(error.Truncated, decode(off[0..n])); + + @memcpy(off[0..n], good[0..n]); + std.mem.writeInt(u16, off[header_len..][0..2], inner + 4, .little); + try testing.expectError(error.Truncated, decode(off[0..n])); + + @memcpy(off[0..n], good[0..n]); + std.mem.writeInt(u16, off[header_len + 2 ..][0..2], inner + 2, .little); + try testing.expectError(error.Truncated, decode(off[0..n])); + + @memcpy(off[0..n], good[0..n]); + std.mem.writeInt(u16, off[header_len + 2 ..][0..2], inner - 1, .little); + try testing.expectError(error.Trailing, decode(off[0..n])); +} + +fn expectTruncatedAtEveryBoundary(full: []const u8) !void { + var scratch: [1024]u8 = undefined; + var n: usize = 0; + while (n < full.len) : (n += 1) { + try testing.expectError(error.Truncated, decode(full[0..n])); + if (n < header_len) continue; + @memcpy(scratch[0..n], full[0..n]); + std.mem.writeInt(u32, scratch[0..4], @intCast(n), .little); + try testing.expectError(error.Truncated, decode(scratch[0..n])); + } + _ = try decode(full); +} + +test "9p: truncation at every field boundary is refused" { + var buf: [512]u8 = undefined; + + try expectTruncatedAtEveryBoundary(try encode( + .{ .tversion = .{ .msize = 8192, .version = "9P2000" } }, + notag, + &buf, + )); + try expectTruncatedAtEveryBoundary(try encode(.{ .twalk = .{ + .fid = 1, + .newfid = 2, + .nwname = 3, + .wname = .{ "usr", "", "bin" } ++ @as([max_welem - 3][]const u8, @splat("")), + } }, 1, &buf)); + try expectTruncatedAtEveryBoundary(try encode( + .{ .tread = .{ .fid = 1, .offset = 0x0102_0304_0506_0708, .count = 8168 } }, + 1, + &buf, + )); + try expectTruncatedAtEveryBoundary(try encode(.{ .rstat = .{ .stat = sample_stat } }, 1, &buf)); + try expectTruncatedAtEveryBoundary(try encode(.{ .rread = .{ .data = "12345678" } }, 1, &buf)); + try expectTruncatedAtEveryBoundary(try encode( + .{ .rwalk = .{ .nwqid = 3, .wqid = @splat(sample_qid) } }, + 1, + &buf, + )); + try expectTruncatedAtEveryBoundary(try encode(.{ .twstat = .{ .fid = 1, .stat = sample_stat } }, 1, &buf)); +} + +test "9p: a size field that disagrees with the buffer is refused" { + var buf: [512]u8 = undefined; + const bytes = try encode(.{ .tclunk = .{ .fid = 1 } }, 1, &buf); + try testing.expectEqual(@as(usize, 11), bytes.len); + + var raw: [64]u8 = undefined; + @memcpy(raw[0..bytes.len], bytes); + + for ([_]u32{ 12, 13, 64, 1 << 20, std.math.maxInt(u32) }) |claim| { + std.mem.writeInt(u32, raw[0..4], claim, .little); + try testing.expectError(error.Truncated, decode(raw[0..bytes.len])); + } + + std.mem.writeInt(u32, raw[0..4], 10, .little); + try testing.expectError(error.Trailing, decode(raw[0..bytes.len])); + + for ([_]u32{ 0, 1, 6 }) |claim| { + std.mem.writeInt(u32, raw[0..4], claim, .little); + try testing.expectError(error.BadValue, decode(raw[0..bytes.len])); + try testing.expectError(error.BadValue, decode(raw[0..header_len])); + } +} + +test "9p: an unknown or illegal type byte is refused" { + var buf: [512]u8 = undefined; + const bytes = try encode(.{ .tclunk = .{ .fid = 1 } }, 1, &buf); + var raw: [64]u8 = undefined; + @memcpy(raw[0..bytes.len], bytes); + + for ([_]u8{ 0, 1, 8, 12, 99, 106, 128, 255 }) |t| { + raw[4] = t; + try testing.expectError(error.BadTag, decode(raw[0..bytes.len])); + } + + var t: u16 = 0; + while (t <= 255) : (t += 1) { + raw[4] = @intCast(t); + const defined = t >= 100 and t <= 127 and t != @intFromEnum(Type.terror); + if (decode(raw[0..bytes.len])) |got| { + try testing.expectEqual(@as(u8, @intCast(t)), @intFromEnum(got.msg.msgType())); + try testing.expect(t == @intFromEnum(Type.tclunk) or + t == @intFromEnum(Type.tremove) or + t == @intFromEnum(Type.tstat) or + t == @intFromEnum(Type.rwrite)); + } else |err| { + if (!defined) try testing.expectEqual(Error.BadTag, err); + } + } +} + +test "9p: trailing bytes inside the size are refused" { + var raw: [64]u8 = undefined; + + var w: Writer = .init(&raw); + try w.putU32(12); + try w.putByte(@intFromEnum(Type.tclunk)); + try w.putU16(1); + try w.putU32(7); + try w.putByte(0xAA); + try testing.expectEqual(@as(usize, 12), w.n); + try testing.expectError(error.Trailing, decode(raw[0..12])); + + w = .init(&raw); + try w.putU32(header_len + 2 + 2); + try w.putByte(@intFromEnum(Type.tflush)); + try w.putU16(1); + try w.putU16(3); + try w.putU16(3); + try testing.expectError(error.Trailing, decode(raw[0..w.n])); +} + +test "9p: frameLen needs four bytes" { + var buf: [512]u8 = undefined; + const bytes = try encode(.{ .tread = .{ .fid = 1, .offset = 0, .count = 8168 } }, 1, &buf); + try testing.expectEqual(@as(usize, 23), bytes.len); + + for (0..4) |n| try testing.expectEqual(@as(?u32, null), frameLen(bytes[0..n])); + try testing.expectEqual(@as(?u32, 23), frameLen(bytes[0..4])); + try testing.expectEqual(@as(?u32, 23), frameLen(bytes)); + + var raw: [4]u8 = .{ 0xFF, 0xFF, 0xFF, 0xFF }; + try testing.expectEqual(@as(?u32, std.math.maxInt(u32)), frameLen(&raw)); + raw = .{ 0, 0, 0, 0 }; + try testing.expectEqual(@as(?u32, 0), frameLen(&raw)); +} + +test "9p: encode refuses a short buffer and writes nothing" { + var buf: [512]u8 = undefined; + const want = (try encode(.{ .rstat = .{ .stat = sample_stat } }, 1, &buf)).len; + + var n: usize = 0; + while (n < want) : (n += 1) { + var scratch: [512]u8 = @splat(0xAA); + try testing.expectError(error.NoSpace, encode(.{ .rstat = .{ .stat = sample_stat } }, 1, scratch[0..n])); + for (scratch) |b| try testing.expectEqual(@as(u8, 0xAA), b); + } + + var exact: [512]u8 = @splat(0xAA); + try testing.expectEqual(want, (try encode(.{ .rstat = .{ .stat = sample_stat } }, 1, exact[0..want])).len); + try testing.expectEqual(@as(u8, 0xAA), exact[want]); +} + +test "9p: byte for byte against u9fs convS2M" { + var buf: [512]u8 = undefined; + + try testing.expectEqualSlices(u8, &.{ + 0x13, 0x00, 0x00, 0x00, // size = 19 + 0x64, // Tversion = 100 + 0xff, 0xff, // NOTAG + 0x00, 0x20, 0x00, 0x00, // msize = 8192 + 0x06, 0x00, // n = 6 + '9', 'P', + '2', '0', + '0', '0', + }, try encode(.{ .tversion = .{ .msize = 8192, .version = "9P2000" } }, notag, &buf)); + + try testing.expectEqualSlices(u8, &.{ + 0x1b, 0x00, 0x00, 0x00, // size = 27 + 0x6e, // Twalk = 110 + 0x01, 0x00, // tag = 1 + 0x01, 0x00, 0x00, 0x00, // fid = 1 + 0x02, 0x00, 0x00, 0x00, // newfid = 2 + 0x02, 0x00, // nwname = 2 + 0x03, 0x00, + 'u', 's', + 'r', 0x03, + 0x00, 'b', + 'i', 'n', + }, try encode(.{ .twalk = .{ + .fid = 1, + .newfid = 2, + .nwname = 2, + .wname = .{ "usr", "bin" } ++ @as([max_welem - 2][]const u8, @splat("")), + } }, 1, &buf)); + + try testing.expectEqualSlices(u8, &.{ + 0x0e, 0x00, 0x00, 0x00, // size = 14 + 0x75, // Rread = 117 + 0x09, 0x00, // tag = 9 + 0x03, 0x00, 0x00, 0x00, // count = 3 + 'a', 'b', 'c', + }, try encode(.{ .rread = .{ .data = "abc" } }, 9, &buf)); + + const one: Stat = .{ + .type = 0, + .dev = 0, + .qid = .{ .type = qtdir, .version = 1, .path = 2 }, + .mode = dmdir | 0o755, + .atime = 3, + .mtime = 4, + .length = 0, + .name = "a", + .uid = "u", + .gid = "g", + .muid = "m", + }; + try testing.expectEqual(@as(u16, 51), try one.size()); + try testing.expectEqualSlices(u8, &.{ + 0x3e, 0x00, 0x00, 0x00, // size = 62 + 0x7d, // Rstat = 125 + 0x07, 0x00, // tag = 7 + 0x35, 0x00, // OUTER count = 53 = 51 + 2 + 0x33, 0x00, // stat size = 51, excluding these two + 0x00, 0x00, // type + 0x00, 0x00, 0x00, 0x00, // dev + 0x80, // qid.type = QTDIR + 0x01, 0x00, 0x00, 0x00, // qid.version = 1 + 0x02, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, // qid.path = 2 + 0xed, 0x01, 0x00, 0x80, // mode = DMDIR | 0755 + 0x03, 0x00, 0x00, 0x00, // atime + 0x04, 0x00, 0x00, 0x00, // mtime + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, // length + 0x01, 0x00, 'a', // name + 0x01, 0x00, 'u', // uid + 0x01, 0x00, 'g', // gid + 0x01, 0x00, 'm', // muid + }, try encode(.{ .rstat = .{ .stat = one } }, 7, &buf)); + + try testing.expectEqual(qtdir, @as(u8, @intCast(dmdir >> 24))); + try testing.expectEqual(qtappend, @as(u8, @intCast(dmappend >> 24))); + try testing.expectEqual(qtexcl, @as(u8, @intCast(dmexcl >> 24))); + try testing.expectEqual(qtauth, @as(u8, @intCast(dmauth >> 24))); + try testing.expectEqual(qttmp, @as(u8, @intCast(dmtmp >> 24))); + try testing.expectEqual(@as(u32, 0o777), dmperm); +} |
