diff options
Diffstat (limited to 'src/client.zig')
| -rw-r--r-- | src/client.zig | 392 |
1 files changed, 392 insertions, 0 deletions
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 }, + } }; + } +}; |
