a73x

src/client/forward.zig

Ref:   Size: 40.9 KiB   History

//! Process-local TCP listener manager for mux forwarding-role connections.
//! Listeners survive transport reconnects; individual accepted streams do not.
const std = @import("std");
const client = @import("client.zig");
const proto = @import("term").protocol;
const Wire = @import("buffered_wire.zig").Wire;
const client_os = @import("client_os");

pub const max_rules: usize = 16;
const transport_queue_max = @import("buffered_wire.zig").transport_queue_max;
const control_reserve: usize = 64 * 1024;
const channel_queue_max: usize = proto.forward_initial_credit;
const handshake_ms: i64 = 2000;
const reconnect_pause_ms: i32 = 100;

pub const Rule = struct {
    local_port: u16,
    remote_port: u16,

    pub fn parse(text: []const u8) !Rule {
        const colon = std.mem.indexOfScalar(u8, text, ':') orelse return error.InvalidForward;
        if (colon == 0 or colon + 1 == text.len or std.mem.indexOfScalarPos(u8, text, colon + 1, ':') != null) return error.InvalidForward;
        for (text[0..colon]) |byte| if (byte < '0' or byte > '9') return error.InvalidForward;
        for (text[colon + 1 ..]) |byte| if (byte < '0' or byte > '9') return error.InvalidForward;
        const local = std.fmt.parseInt(u16, text[0..colon], 10) catch return error.InvalidForward;
        const remote = std.fmt.parseInt(u16, text[colon + 1 ..], 10) catch return error.InvalidForward;
        if (local == 0 or remote == 0) return error.InvalidForward;
        return .{ .local_port = local, .remote_port = remote };
    }

    pub fn eql(a: Rule, b: Rule) bool {
        return a.local_port == b.local_port and a.remote_port == b.remote_port;
    }
};

const Listener = struct { fd: std.posix.fd_t, rule: Rule };
const Channel = struct {
    id: u32,
    fd: std.posix.fd_t,
    opening: bool = true,
    send_credit: u32 = 0,
    recv_credit: u32 = 0,
    to_local: std.ArrayList(u8) = .empty,
    local_eof: bool = false,
    peer_eof: bool = false,
    write_shutdown: bool = false,
};

pub const Manager = struct {
    alloc: std.mem.Allocator,
    arena: std.heap.ArenaAllocator,
    target: client.Target,
    listeners: [max_rules]?Listener = @splat(null),
    listener_len: usize = 0,
    channels: [proto.forward_channels_max]?Channel = @splat(null),
    next_id: u32 = 1,
    wake_pipe: [2]std.posix.fd_t,
    closing: std.atomic.Value(bool) = .init(false),
    thread: ?std.Thread = null,

    /// Binds every unique listener before returning, so an explicit collision
    /// is a synchronous CLI failure rather than a background warning.
    pub fn init(alloc: std.mem.Allocator, source: client.Target, rules: []const Rule) !*Manager {
        if (rules.len == 0) return error.InvalidForward;
        const self = try alloc.create(Manager);
        errdefer alloc.destroy(self);
        var arena = std.heap.ArenaAllocator.init(alloc);
        errdefer arena.deinit();
        var target = try client.discovery.cloneTarget(arena.allocator(), source);
        if (target == .hand) {
            target.hand.asked = false;
            target.hand.narrate = false;
        }
        const wake = try std.posix.pipe2(.{ .NONBLOCK = true, .CLOEXEC = true });
        errdefer {
            std.posix.close(wake[0]);
            std.posix.close(wake[1]);
        }
        self.* = .{ .alloc = alloc, .arena = arena, .target = target, .wake_pipe = wake };
        errdefer self.closeListeners();
        for (rules) |rule| {
            var duplicate = false;
            for (self.listeners[0..self.listener_len]) |slot| if (slot != null and Rule.eql(slot.?.rule, rule)) {
                duplicate = true;
                break;
            };
            if (duplicate) continue;
            if (self.listener_len == max_rules) return error.TooManyForwards;
            // Two different destinations cannot own the same local port.
            for (self.listeners[0..self.listener_len]) |slot| if (slot != null and slot.?.rule.local_port == rule.local_port)
                return error.ForwardConflict;
            self.listeners[self.listener_len] = .{ .fd = try bindLoopback(rule.local_port), .rule = rule };
            self.listener_len += 1;
        }
        return self;
    }

    pub fn start(self: *Manager) !void {
        if (self.thread != null) return;
        self.thread = try std.Thread.spawn(.{}, entry, .{self});
    }

    pub fn stop(self: *Manager) void {
        if (self.closing.swap(true, .acq_rel)) return;
        _ = std.posix.write(self.wake_pipe[1], &.{1}) catch {};
        if (self.thread) |thread| thread.join();
        self.dropChannels(null);
        self.closeListeners();
        std.posix.close(self.wake_pipe[0]);
        std.posix.close(self.wake_pipe[1]);
        self.arena.deinit();
        const alloc = self.alloc;
        alloc.destroy(self);
    }

    pub fn matches(self: *const Manager, target: client.Target) bool {
        return targetEqual(self.target, target);
    }

    fn closeListeners(self: *Manager) void {
        for (&self.listeners) |*slot| if (slot.*) |listener| {
            std.posix.close(listener.fd);
            slot.* = null;
        };
        self.listener_len = 0;
    }

    fn entry(self: *Manager) void {
        var warned = false;
        while (!self.closing.load(.acquire)) {
            var job = DialJob.init(self.alloc, self.target) catch return;
            defer job.deinit();
            job.start() catch return;
            while (!job.done.load(.acquire) and !self.closing.load(.acquire)) self.refusePendingAccepts(50);
            if (self.closing.load(.acquire)) job.cancel();
            job.join();
            var tr = job.take() orelse {
                if (!self.pauseAfterFailure()) return;
                continue;
            };
            defer tr.close();
            var wire = Wire.init(self.alloc, &tr) catch {
                if (!self.pauseAfterFailure()) return;
                continue;
            };
            defer wire.deinit();
            if (!self.handshake(&wire)) {
                if (!self.closing.load(.acquire) and !warned) {
                    std.debug.print("muxg: daemon does not support native port forwarding\n", .{});
                    warned = true;
                }
                self.dropChannels(null);
                if (!self.pauseAfterFailure()) return;
                continue;
            }
            warned = false;
            self.connected(&wire);
            self.dropChannels(null);
            if (!self.closing.load(.acquire) and !self.pauseAfterFailure()) return;
        }
    }

    fn handshake(self: *Manager, wire: *Wire) bool {
        const hello = proto.encodeForwardHello();
        wire.send(.forward_hello, &hello) catch return false;
        const until = std.time.milliTimestamp() + handshake_ms;
        while (!self.closing.load(.acquire) and std.time.milliTimestamp() < until) {
            wire.tr.service();
            wire.flush() catch return false;
            switch (wire.readLimited(transport_queue_max) catch return false) {
                .frame => |frame| {
                    defer frame.deinit(self.alloc);
                    if (frame.type != .forward_ready) return false;
                    const version = proto.decodeForwardHello(frame.payload) catch return false;
                    return version == proto.forward_version;
                },
                .closed => return false,
                .incomplete => {},
            }
            var fds = [_]std.posix.pollfd{
                .{ .fd = wire.tr.pollFd(), .events = std.posix.POLL.IN, .revents = 0 },
                .{ .fd = wire.pendingWriteFd() orelse -1, .events = std.posix.POLL.OUT, .revents = 0 },
                .{ .fd = self.wake_pipe[0], .events = std.posix.POLL.IN, .revents = 0 },
                .{ .fd = wire.tr.errFd() orelse -1, .events = std.posix.POLL.IN, .revents = 0 },
            };
            _ = std.posix.poll(&fds, wire.tr.timeoutMs(50)) catch return false;
            if (fds[3].revents != 0) wire.tr.drainErr();
            self.refuseReadyAccepts();
        }
        return false;
    }

    fn connected(self: *Manager, wire: *Wire) void {
        while (!self.closing.load(.acquire)) {
            wire.tr.service();
            wire.flush() catch return;
            if (wire.pendingBytes() > transport_queue_max) return;
            var fds: [4 + max_rules + proto.forward_channels_max]std.posix.pollfd = undefined;
            fds[0] = .{ .fd = wire.tr.pollFd(), .events = std.posix.POLL.IN, .revents = 0 };
            fds[1] = .{ .fd = wire.pendingWriteFd() orelse -1, .events = std.posix.POLL.OUT, .revents = 0 };
            fds[2] = .{ .fd = self.wake_pipe[0], .events = std.posix.POLL.IN, .revents = 0 };
            fds[3] = .{ .fd = wire.tr.errFd() orelse -1, .events = std.posix.POLL.IN, .revents = 0 };
            const listener_base = 4;
            for (0..max_rules) |i| fds[listener_base + i] = .{ .fd = if (self.listeners[i]) |l| l.fd else -1, .events = std.posix.POLL.IN, .revents = 0 };
            const chan_base = listener_base + max_rules;
            for (0..proto.forward_channels_max) |i| {
                if (self.channels[i]) |ch| {
                    var events: i16 = 0;
                    if (!ch.opening and !ch.local_eof and self.dataReadLimit(wire, ch) != 0) events |= std.posix.POLL.IN;
                    if (ch.to_local.items.len != 0 or (ch.peer_eof and !ch.write_shutdown)) events |= std.posix.POLL.OUT;
                    fds[chan_base + i] = .{ .fd = if (events == 0) -1 else ch.fd, .events = events, .revents = 0 };
                } else fds[chan_base + i] = .{ .fd = -1, .events = 0, .revents = 0 };
            }
            _ = std.posix.poll(&fds, wire.tr.timeoutMs(50)) catch return;
            if (fds[2].revents != 0) return;
            if (fds[3].revents != 0) wire.tr.drainErr();
            if (fds[0].revents != 0 or wire.tr.link == .quic) if (!self.receive(wire)) return;
            if (fds[1].revents != 0) wire.flush() catch return;
            for (0..max_rules) |i| if (fds[listener_base + i].revents != 0 and !self.acceptOne(wire, i)) return;
            for (0..proto.forward_channels_max) |i| {
                if (self.channels[i] == null or fds[chan_base + i].revents == 0) continue;
                const revents = fds[chan_base + i].revents;
                if (revents & std.posix.POLL.OUT != 0 and !self.flushLocal(wire, i)) return;
                if (self.channels[i] != null and revents & ~@as(i16, std.posix.POLL.OUT) != 0 and !self.readLocal(wire, i)) return;
            }
        }
    }

    fn receive(self: *Manager, wire: *Wire) bool {
        for (0..64) |_| switch (wire.readLimited(transport_queue_max) catch return false) {
            .closed => return false,
            .incomplete => return true,
            .frame => |frame| {
                defer frame.deinit(self.alloc);
                if (!self.handleFrame(wire, frame)) return false;
            },
        };
        return true;
    }

    fn handleFrame(self: *Manager, wire: *Wire, frame: proto.Frame) bool {
        switch (frame.type) {
            .forward_open_result => {
                const result = proto.decodeForwardOpenResult(frame.payload) catch return false;
                const ci = self.findChannel(result.id) orelse return true;
                if (!result.ok) {
                    _ = self.dropChannel(ci, false, wire);
                    return true;
                }
                const ch = &self.channels[ci].?;
                ch.opening = false;
                ch.recv_credit = proto.forward_initial_credit;
                const credit = proto.encodeForwardCredit(.{ .id = ch.id, .amount = proto.forward_initial_credit });
                self.send(wire, .forward_credit, &credit) catch return false;
            },
            .forward_credit => {
                const credit = proto.decodeForwardCredit(frame.payload) catch return false;
                const ci = self.findChannel(credit.id) orelse return true;
                const ch = &self.channels[ci].?;
                if (credit.amount > proto.forward_initial_credit -| ch.send_credit) return false;
                ch.send_credit += credit.amount;
            },
            .forward_data => {
                if (proto.forwardDataOversize(frame.payload)) return false;
                const id = proto.decodeForwardId(frame.payload) catch return false;
                const ci = self.findChannel(id) orelse return true;
                const ch = &self.channels[ci].?;
                const data = frame.payload[proto.forward_id_len..];
                if (ch.peer_eof or data.len > ch.recv_credit or ch.to_local.items.len + data.len > channel_queue_max) {
                    return self.dropChannel(ci, true, wire);
                }
                ch.recv_credit -= @intCast(data.len);
                ch.to_local.appendSlice(self.alloc, data) catch return self.dropChannel(ci, true, wire);
                if (self.channels[ci] != null and !self.flushLocal(wire, ci)) return false;
            },
            .forward_half_close => {
                const id = proto.decodeForwardId(frame.payload) catch return false;
                if (frame.payload.len != proto.forward_id_len) return false;
                const ci = self.findChannel(id) orelse return true;
                self.channels[ci].?.peer_eof = true;
                if (!self.flushLocal(wire, ci)) return false;
            },
            .forward_reset => {
                const id = proto.decodeForwardId(frame.payload) catch return false;
                if (frame.payload.len != proto.forward_id_len) return false;
                if (self.findChannel(id)) |ci| _ = self.dropChannel(ci, false, wire);
            },
            else => return false,
        }
        return true;
    }

    fn acceptOne(self: *Manager, wire: *Wire, li: usize) bool {
        const listener = self.listeners[li] orelse return true;
        const fd = std.posix.accept(listener.fd, null, null, std.posix.SOCK.NONBLOCK | std.posix.SOCK.CLOEXEC) catch return true;
        var ci: ?usize = null;
        for (self.channels, 0..) |slot, i| if (slot == null) {
            ci = i;
            break;
        };
        const at = ci orelse {
            std.posix.close(fd);
            return true;
        };
        const id = self.next_id;
        if (id == std.math.maxInt(u32)) {
            std.posix.close(fd);
            return true;
        }
        self.next_id += 1;
        self.channels[at] = .{ .id = id, .fd = fd };
        const payload = proto.encodeForwardOpen(.{ .id = id, .port = listener.rule.remote_port });
        self.send(wire, .forward_open, &payload) catch {
            _ = self.dropChannel(at, false, wire);
            return false;
        };
        return true;
    }

    fn readLocal(self: *Manager, wire: *Wire, ci: usize) bool {
        const ch = &self.channels[ci].?;
        if (ch.local_eof or ch.send_credit == 0) return true;
        // Every descriptor in a poll snapshot may have observed the same
        // shared headroom. Recompute immediately before consuming TCP bytes.
        const want = self.dataReadLimit(wire, ch.*);
        if (want == 0) return true;
        var buf: [proto.forward_data_max]u8 = undefined;
        const n = std.posix.read(ch.fd, buf[0..want]) catch |err| switch (err) {
            error.WouldBlock => return true,
            else => return self.dropChannel(ci, true, wire),
        };
        if (n == 0) {
            ch.local_eof = true;
            const id = proto.encodeForwardId(ch.id);
            self.send(wire, .forward_half_close, &id) catch return false;
            if (self.channels[ci] == null) return true;
            return self.maybeFinish(ci, wire);
        }
        ch.send_credit -= @intCast(n);
        var payload: [proto.forward_id_len + proto.forward_data_max]u8 = undefined;
        std.mem.writeInt(u32, payload[0..4], ch.id, .little);
        @memcpy(payload[4..][0..n], buf[0..n]);
        self.send(wire, .forward_data, payload[0 .. 4 + n]) catch return false;
        return true;
    }

    fn flushLocal(self: *Manager, wire: *Wire, ci: usize) bool {
        const ch = &self.channels[ci].?;
        if (ch.to_local.items.len != 0) {
            const n = client_os.sendNoSigNoWait(ch.fd, ch.to_local.items) catch |err| switch (err) {
                error.WouldBlock => return true,
                else => return self.dropChannel(ci, true, wire),
            };
            ch.to_local.replaceRangeAssumeCapacity(0, n, &.{});
            ch.recv_credit += @intCast(n);
            const credit = proto.encodeForwardCredit(.{ .id = ch.id, .amount = @intCast(n) });
            self.send(wire, .forward_credit, &credit) catch return false;
        }
        if (self.channels[ci] == null) return true;
        const current = &self.channels[ci].?;
        if (current.peer_eof and current.to_local.items.len == 0 and !current.write_shutdown) {
            std.posix.shutdown(current.fd, .send) catch {};
            current.write_shutdown = true;
        }
        return self.maybeFinish(ci, wire);
    }

    fn maybeFinish(self: *Manager, ci: usize, wire: *Wire) bool {
        if (self.channels[ci] == null) return true;
        const ch = &self.channels[ci].?;
        if (ch.local_eof and ch.peer_eof and ch.to_local.items.len == 0) return self.dropChannel(ci, false, wire);
        return true;
    }

    fn findChannel(self: *const Manager, id: u32) ?usize {
        for (self.channels, 0..) |ch, i| if (ch != null and ch.?.id == id) return i;
        return null;
    }

    fn dropChannel(self: *Manager, ci: usize, notify: bool, wire: ?*Wire) bool {
        if (self.channels[ci]) |*ch| {
            const id = ch.id;
            std.posix.close(ch.fd);
            ch.to_local.deinit(self.alloc);
            self.channels[ci] = null;
            if (notify) if (wire) |w| {
                const payload = proto.encodeForwardId(id);
                self.send(w, .forward_reset, &payload) catch return false;
            };
        }
        return true;
    }

    fn dropChannels(self: *Manager, wire: ?*Wire) void {
        for (0..proto.forward_channels_max) |i| _ = self.dropChannel(i, false, wire);
    }

    fn send(_: *Manager, wire: *Wire, kind: proto.MsgType, payload: []const u8) !void {
        if (wire.pendingBytes() + proto.frame_header_len + payload.len > transport_queue_max) return error.ForwardQueueFull;
        try wire.send(kind, payload);
    }

    fn dataReadLimit(_: *const Manager, wire: *const Wire, ch: Channel) usize {
        const overhead = proto.frame_header_len + proto.forward_id_len;
        const pending = wire.pendingBytes();
        if (pending >= transport_queue_max - control_reserve - overhead) return 0;
        const room = transport_queue_max - control_reserve - overhead - pending;
        return @min(proto.forward_data_max, @min(@as(usize, ch.send_credit), room));
    }

    /// Every failed setup, including a peer that hangs up immediately after
    /// dialling, waits through this interruptible pause. This prevents a
    /// rejected peer from spinning a reconnect loop while listeners continue
    /// to refuse new local accepts and stop still wakes promptly.
    fn pauseAfterFailure(self: *Manager) bool {
        // Listener readability is expected while a transport is down. It may
        // only refuse bounded work; it must not restart this retry deadline.
        const deadline = monotonicMs() + reconnect_pause_ms;
        while (!self.closing.load(.acquire)) {
            const remaining = deadline - monotonicMs();
            if (remaining <= 0) return true;
            var fds: [1 + max_rules]std.posix.pollfd = undefined;
            fds[0] = .{ .fd = self.wake_pipe[0], .events = std.posix.POLL.IN, .revents = 0 };
            for (0..max_rules) |i| fds[1 + i] = .{ .fd = if (self.listeners[i]) |l| l.fd else -1, .events = std.posix.POLL.IN, .revents = 0 };
            _ = std.posix.poll(&fds, @intCast(@min(remaining, @as(i64, std.math.maxInt(i32))))) catch return false;
            if (fds[0].revents != 0 or self.closing.load(.acquire)) return false;
            // At most one accept per wake: a flood cannot consume the pause.
            for (0..max_rules) |i| if (fds[1 + i].revents != 0) {
                const listener = self.listeners[i] orelse continue;
                const fd = std.posix.accept(listener.fd, null, null, std.posix.SOCK.NONBLOCK | std.posix.SOCK.CLOEXEC) catch continue;
                std.posix.close(fd);
                break;
            };
        }
        return false;
    }

    fn refusePendingAccepts(self: *Manager, timeout_ms: i32) void {
        var fds: [1 + max_rules]std.posix.pollfd = undefined;
        fds[0] = .{ .fd = self.wake_pipe[0], .events = std.posix.POLL.IN, .revents = 0 };
        for (0..max_rules) |i| fds[1 + i] = .{ .fd = if (self.listeners[i]) |l| l.fd else -1, .events = std.posix.POLL.IN, .revents = 0 };
        _ = std.posix.poll(&fds, timeout_ms) catch return;
        for (0..max_rules) |i| if (fds[1 + i].revents != 0) {
            const listener = self.listeners[i] orelse continue;
            const fd = std.posix.accept(listener.fd, null, null, std.posix.SOCK.NONBLOCK | std.posix.SOCK.CLOEXEC) catch continue;
            std.posix.close(fd);
        };
    }

    fn refuseReadyAccepts(self: *Manager) void {
        self.refusePendingAccepts(0);
    }
};

const DialJob = struct {
    alloc: std.mem.Allocator,
    target: client.Target,
    cancel_pipe: [2]std.posix.fd_t,
    done: std.atomic.Value(bool) = .init(false),
    thread: ?std.Thread = null,
    transport: ?client.Transport = null,

    fn init(alloc: std.mem.Allocator, target: client.Target) !DialJob {
        return .{ .alloc = alloc, .target = target, .cancel_pipe = try std.posix.pipe2(.{ .NONBLOCK = true, .CLOEXEC = true }) };
    }
    fn deinit(self: *DialJob) void {
        if (self.transport) |*tr| tr.close();
        std.posix.close(self.cancel_pipe[0]);
        std.posix.close(self.cancel_pipe[1]);
    }
    fn start(self: *DialJob) !void {
        self.thread = try std.Thread.spawn(.{}, run, .{self});
    }
    fn run(self: *DialJob) void {
        defer self.done.store(true, .release);
        var dial: client.handoff.Dial = .{};
        self.transport = client.Transport.openBounded(self.alloc, self.target, null, self.cancel_pipe[0], &dial, transport_queue_max) catch null;
    }
    fn cancel(self: *DialJob) void {
        _ = std.posix.write(self.cancel_pipe[1], &.{client.interrupt.detach_key}) catch {};
    }
    fn join(self: *DialJob) void {
        if (self.thread) |thread| thread.join();
        self.thread = null;
    }
    fn take(self: *DialJob) ?client.Transport {
        const value = self.transport;
        self.transport = null;
        return value;
    }
};

fn monotonicMs() i64 {
    const t = std.posix.clock_gettime(.MONOTONIC) catch return std.time.milliTimestamp();
    return @as(i64, t.sec) * 1000 + @divFloor(t.nsec, 1_000_000);
}

fn bindLoopback(port: u16) !std.posix.fd_t {
    const fd = try std.posix.socket(std.posix.AF.INET, std.posix.SOCK.STREAM | std.posix.SOCK.NONBLOCK | std.posix.SOCK.CLOEXEC, std.posix.IPPROTO.TCP);
    errdefer std.posix.close(fd);
    const addr = try std.net.Address.parseIp4("127.0.0.1", port);
    try std.posix.bind(fd, &addr.any, addr.getOsSockLen());
    try std.posix.listen(fd, 128);
    return fd;
}

pub fn targetEqual(a: client.Target, b: client.Target) bool {
    if (std.meta.activeTag(a) != std.meta.activeTag(b)) return false;
    return switch (a) {
        .sock => |x| std.mem.eql(u8, x, b.sock),
        .via => |x| std.mem.eql(u8, x, b.via),
        .quic => |x| std.mem.eql(u8, x.host_port, b.quic.host_port) and std.mem.eql(u8, x.key_path, b.quic.key_path),
        .hand => |x| std.mem.eql(u8, x.host, b.hand.host) and argvEqual(x.ssh_argv, b.hand.ssh_argv),
    };
}

fn argvEqual(a: []const []const u8, b: []const []const u8) bool {
    if (a.len != b.len) return false;
    for (a, b) |x, y| if (!std.mem.eql(u8, x, y)) return false;
    return true;
}

test "forward rule accepts only nonzero decimal port pairs" {
    try std.testing.expectEqual(Rule{ .local_port = 8080, .remote_port = 80 }, try Rule.parse("8080:80"));
    for ([_][]const u8{ "", "80", "0:1", "1:0", "x:2", "+1:2", "1: 2", "1:2:3", "65536:2" }) |bad|
        try std.testing.expectError(error.InvalidForward, Rule.parse(bad));
}

fn testUnusedPort() !u16 {
    const fd = try bindLoopback(0);
    defer std.posix.close(fd);
    var actual: std.posix.sockaddr.storage = undefined;
    var len: std.posix.socklen_t = @sizeOf(@TypeOf(actual));
    try std.posix.getsockname(fd, @ptrCast(&actual), &len);
    return std.net.Address.initPosix(@ptrCast(@alignCast(&actual))).getPort();
}

test "manager deduplicates before enforcing its unique listener cap" {
    const rule: Rule = .{ .local_port = try testUnusedPort(), .remote_port = 80 };
    const duplicates: [max_rules + 3]Rule = @splat(rule);
    const manager = try Manager.init(std.testing.allocator, .{ .sock = "/not-dialled" }, &duplicates);
    defer manager.stop();
    try std.testing.expectEqual(@as(usize, 1), manager.listener_len);
}

test "manager rejects conflicting destinations without leaking its first bind" {
    const port = try testUnusedPort();
    try std.testing.expectError(
        error.ForwardConflict,
        Manager.init(std.testing.allocator, .{ .sock = "/not-dialled" }, &.{
            .{ .local_port = port, .remote_port = 80 },
            .{ .local_port = port, .remote_port = 81 },
        }),
    );
    const rebound = try bindLoopback(port);
    std.posix.close(rebound);
}

test "manager keeps the listener reserved and refuses accepts while unavailable" {
    const port = try testUnusedPort();
    const manager = try Manager.init(std.testing.allocator, .{ .sock = "/not-dialled" }, &.{.{ .local_port = port, .remote_port = 80 }});
    defer manager.stop();

    const fd = try std.posix.socket(std.posix.AF.INET, std.posix.SOCK.STREAM | std.posix.SOCK.CLOEXEC, std.posix.IPPROTO.TCP);
    defer std.posix.close(fd);
    const addr = try std.net.Address.parseIp4("127.0.0.1", port);
    try std.posix.connect(fd, &addr.any, addr.getOsSockLen());
    manager.refusePendingAccepts(100);
    var pfd = [_]std.posix.pollfd{.{ .fd = fd, .events = std.posix.POLL.IN, .revents = 0 }};
    try std.testing.expectEqual(@as(usize, 1), try std.posix.poll(&pfd, 1000));
    try std.testing.expect(pfd[0].revents & (std.posix.POLL.IN | std.posix.POLL.HUP | std.posix.POLL.ERR) != 0);
}

test "target identity scopes forwarding independently of session and handoff narration" {
    try std.testing.expect(targetEqual(.{ .sock = "/a" }, .{ .sock = "/a" }));
    try std.testing.expect(!targetEqual(.{ .sock = "/a" }, .{ .sock = "/b" }));
    try std.testing.expect(!targetEqual(.{ .sock = "/a" }, .{ .via = "/a" }));
    try std.testing.expect(targetEqual(
        .{ .quic = .{ .host_port = "host:1", .key_path = "/key", .idle_ms = 1 } },
        .{ .quic = .{ .host_port = "host:1", .key_path = "/key", .idle_ms = 999 } },
    ));
    try std.testing.expect(targetEqual(
        .{ .hand = .{ .host = "host", .ssh_argv = &.{ "ssh", "host" }, .cache_path = null, .asked = true, .narrate = true } },
        .{ .hand = .{ .host = "host", .ssh_argv = &.{ "ssh", "host" }, .cache_path = "/elsewhere" } },
    ));
}

fn testSocketPair() ![2]std.posix.fd_t {
    var pair: [2]std.posix.fd_t = undefined;
    if (std.c.socketpair(std.posix.AF.UNIX, std.posix.SOCK.STREAM, 0, &pair) != 0) return error.SocketPairFailed;
    return pair;
}

fn fillSocket(fd: std.posix.fd_t) !usize {
    var buf: [4096]u8 = @splat(0xa5);
    var total: usize = 0;
    while (true) total += std.posix.write(fd, &buf) catch |err| switch (err) {
        error.WouldBlock => return total,
        else => return err,
    };
}

const NoisyHelper = struct {
    err_fd: std.posix.fd_t,
    transport_read_fd: std.posix.fd_t = -1,
    transport_write_fd: std.posix.fd_t = -1,
    wake_fd: std.posix.fd_t = -1,
    completed: *std.atomic.Value(bool),

    const noise_bytes = 128 * 1024;

    fn floodErr(self: NoisyHelper) bool {
        const flags = std.posix.fcntl(self.err_fd, std.posix.F.GETFL, 0) catch return false;
        const bits: u32 = @bitCast(std.posix.O{ .NONBLOCK = true });
        _ = std.posix.fcntl(self.err_fd, std.posix.F.SETFL, flags | bits) catch return false;
        const until = std.time.milliTimestamp() + handshake_ms;
        var buf: [4096]u8 = @splat('x');
        var written: usize = 0;
        while (written < noise_bytes) {
            const n = std.posix.write(self.err_fd, buf[0..@min(buf.len, noise_bytes - written)]) catch |err| switch (err) {
                error.WouldBlock => {
                    if (std.time.milliTimestamp() >= until) return false;
                    var fds = [_]std.posix.pollfd{.{ .fd = self.err_fd, .events = std.posix.POLL.OUT, .revents = 0 }};
                    _ = std.posix.poll(&fds, 10) catch return false;
                    continue;
                },
                else => return false,
            };
            if (n == 0) return false;
            written += n;
        }
        return true;
    }

    fn answerHandshake(self: NoisyHelper) void {
        if (!self.floodErr()) return;
        var frame = (proto.readFrame(std.testing.allocator, self.transport_read_fd) catch return) orelse return;
        defer frame.deinit(std.testing.allocator);
        if (frame.type != .forward_hello) return;
        const ready = proto.encodeForwardHello();
        proto.writeFrame(self.transport_write_fd, .forward_ready, &ready) catch return;
        self.completed.store(true, .release);
    }

    fn wakeConnected(self: NoisyHelper) void {
        self.completed.store(self.floodErr(), .release);
        _ = std.posix.write(self.wake_fd, &.{1}) catch {};
    }
};

fn testPipe() ![2]std.posix.fd_t {
    return std.posix.pipe2(.{ .CLOEXEC = true });
}

test "forward handshake drains noisy helper stderr" {
    const alloc = std.testing.allocator;
    const port = try testUnusedPort();
    const manager = try Manager.init(alloc, .{ .sock = "/not-dialled" }, &.{.{ .local_port = port, .remote_port = 80 }});
    defer manager.stop();

    const incoming = try testPipe();
    defer std.posix.close(incoming[0]);
    defer std.posix.close(incoming[1]);
    const outgoing = try testPipe();
    defer std.posix.close(outgoing[0]);
    defer std.posix.close(outgoing[1]);
    const err_pipe = try testPipe();
    defer std.posix.close(err_pipe[0]);
    defer std.posix.close(err_pipe[1]);

    var tr: client.Transport = .{
        .link = .{ .pipe = .{
            .child = std.process.Child.init(&.{"unused"}, alloc),
            .r = incoming[0],
            .w = outgoing[1],
        } },
        .err_fd = err_pipe[0],
    };
    var wire = try Wire.init(alloc, &tr);
    defer wire.deinit();

    var completed: std.atomic.Value(bool) = .init(false);
    const helper = try std.Thread.spawn(.{}, NoisyHelper.answerHandshake, .{NoisyHelper{
        .err_fd = err_pipe[1],
        .transport_read_fd = outgoing[0],
        .transport_write_fd = incoming[1],
        .completed = &completed,
    }});

    const handshake_ok = manager.handshake(&wire);
    helper.join();
    try std.testing.expect(handshake_ok);
    try std.testing.expect(completed.load(.acquire));
}

test "connected forwarding drains noisy helper stderr" {
    const alloc = std.testing.allocator;
    const port = try testUnusedPort();
    const manager = try Manager.init(alloc, .{ .sock = "/not-dialled" }, &.{.{ .local_port = port, .remote_port = 80 }});
    defer manager.stop();

    const incoming = try testPipe();
    defer std.posix.close(incoming[0]);
    defer std.posix.close(incoming[1]);
    const outgoing = try testPipe();
    defer std.posix.close(outgoing[0]);
    defer std.posix.close(outgoing[1]);
    const err_pipe = try testPipe();
    defer std.posix.close(err_pipe[0]);
    defer std.posix.close(err_pipe[1]);

    var tr: client.Transport = .{
        .link = .{ .pipe = .{
            .child = std.process.Child.init(&.{"unused"}, alloc),
            .r = incoming[0],
            .w = outgoing[1],
        } },
        .err_fd = err_pipe[0],
    };
    var wire = try Wire.init(alloc, &tr);
    defer wire.deinit();

    var completed: std.atomic.Value(bool) = .init(false);
    const helper = try std.Thread.spawn(.{}, NoisyHelper.wakeConnected, .{NoisyHelper{
        .err_fd = err_pipe[1],
        .wake_fd = manager.wake_pipe[1],
        .completed = &completed,
    }});

    manager.connected(&wire);
    helper.join();
    try std.testing.expect(completed.load(.acquire));
}

fn drainQueuedPrefix(wire: *Wire, fd: std.posix.fd_t, count: usize) !void {
    var left = count;
    var buf: [64 * 1024]u8 = undefined;
    while (left != 0) {
        try wire.flush();
        const n = try std.posix.read(fd, buf[0..@min(buf.len, left)]);
        if (n == 0) return error.UnexpectedEof;
        left -= n;
    }
}

test "manager stop interrupts an opening dial blocked by a full Unix backlog" {
    if (@import("builtin").os.tag != .linux) return error.SkipZigTest;
    const alloc = std.testing.allocator;
    var tmp = try @import("testtmp").TmpDir.make();
    defer tmp.cleanup();
    const path = try std.fmt.allocPrint(alloc, "{s}/full.sock", .{tmp.path()});
    defer alloc.free(path);
    const addr = try std.net.Address.initUnix(path);
    var server = try addr.listen(.{ .kernel_backlog = 1 });
    defer server.deinit();
    var queued: std.ArrayList(std.posix.fd_t) = .empty;
    defer {
        for (queued.items) |fd| std.posix.close(fd);
        queued.deinit(alloc);
    }
    var full = false;
    for (0..16) |_| {
        const fd = try std.posix.socket(std.posix.AF.UNIX, std.posix.SOCK.STREAM | std.posix.SOCK.NONBLOCK | std.posix.SOCK.CLOEXEC, 0);
        const err = std.posix.errno(std.posix.system.connect(fd, &addr.any, addr.getOsSockLen()));
        if (err == .AGAIN) {
            std.posix.close(fd);
            full = true;
            break;
        }
        if (err != .SUCCESS) {
            std.posix.close(fd);
            return error.UnexpectedConnect;
        }
        try queued.append(alloc, fd);
    }
    try std.testing.expect(full);

    const Releaser = struct {
        fn run(listener: *std.net.Server) void {
            std.Thread.sleep(800 * std.time.ns_per_ms);
            const accepted = listener.accept() catch return;
            accepted.stream.close();
        }
    };
    const releaser = try std.Thread.spawn(.{}, Releaser.run, .{&server});
    defer releaser.join();

    const port = try testUnusedPort();
    const manager = try Manager.init(alloc, .{ .sock = path }, &.{.{ .local_port = port, .remote_port = 80 }});
    try manager.start();
    std.Thread.sleep(50 * std.time.ns_per_ms);
    const started = std.time.milliTimestamp();
    manager.stop();
    try std.testing.expect(std.time.milliTimestamp() - started < 500);
}

const RejectingForwardPeer = struct {
    listener: std.net.Server,
    ready_then_closed: std.atomic.Value(u32) = .init(0),
    until_ms: i64,

    fn serve(self: *RejectingForwardPeer) void {
        while (std.time.milliTimestamp() < self.until_ms) {
            const conn = self.listener.accept() catch |err| switch (err) {
                error.WouldBlock => {
                    std.Thread.sleep(std.time.ns_per_ms);
                    continue;
                },
                else => return,
            };
            const frame = (proto.readFrame(std.testing.allocator, conn.stream.handle) catch {
                conn.stream.close();
                continue;
            }) orelse {
                conn.stream.close();
                continue;
            };
            if (frame.type == .forward_hello) {
                const ready = proto.encodeForwardHello();
                proto.writeFrame(conn.stream.handle, .forward_ready, &ready) catch {
                    frame.deinit(std.testing.allocator);
                    conn.stream.close();
                    continue;
                };
                _ = self.ready_then_closed.fetchAdd(1, .acq_rel);
            }
            frame.deinit(std.testing.allocator);
            conn.stream.close();
        }
    }
};

const ListenerFlood = struct {
    port: u16,
    stop: std.atomic.Value(bool) = .init(false),

    fn run(self: *ListenerFlood) void {
        const addr = std.net.Address.parseIp4("127.0.0.1", self.port) catch return;
        while (!self.stop.load(.acquire)) {
            const fd = std.posix.socket(std.posix.AF.INET, std.posix.SOCK.STREAM | std.posix.SOCK.NONBLOCK | std.posix.SOCK.CLOEXEC, std.posix.IPPROTO.TCP) catch continue;
            _ = std.posix.connect(fd, &addr.any, addr.getOsSockLen()) catch {};
            std.posix.close(fd);
        }
    }
};

test "manager keeps retry deadline under continuously readable listeners and stop wakes it" {
    const alloc = std.testing.allocator;
    var tmp = try @import("testtmp").TmpDir.make();
    defer tmp.cleanup();
    const path = try std.fmt.allocPrint(alloc, "{s}/reject.sock", .{tmp.path()});
    defer alloc.free(path);
    const addr = try std.net.Address.initUnix(path);
    var peer = RejectingForwardPeer{
        .listener = try addr.listen(.{}),
        .until_ms = std.time.milliTimestamp() + 450,
    };
    defer peer.listener.deinit();
    const flags = try std.posix.fcntl(peer.listener.stream.handle, std.posix.F.GETFL, 0);
    const nonblock: u32 = @bitCast(std.posix.O{ .NONBLOCK = true });
    _ = try std.posix.fcntl(peer.listener.stream.handle, std.posix.F.SETFL, flags | nonblock);
    const peer_thread = try std.Thread.spawn(.{}, RejectingForwardPeer.serve, .{&peer});
    defer peer_thread.join();

    const port = try testUnusedPort();
    const manager = try Manager.init(alloc, .{ .sock = path }, &.{.{ .local_port = port, .remote_port = 80 }});
    try manager.start();
    var flood = ListenerFlood{ .port = port };
    const flood_thread = try std.Thread.spawn(.{}, ListenerFlood.run, .{&flood});
    defer {
        flood.stop.store(true, .release);
        flood_thread.join();
    }
    std.Thread.sleep(350 * std.time.ns_per_ms);
    const attempts = peer.ready_then_closed.load(.acquire);
    // The first attempt plus fixed 100 ms pauses admits several retries, but
    // not a hot loop even while the local listener is continuously readable.
    try std.testing.expect(attempts >= 2);
    try std.testing.expect(attempts <= 5);

    const started = std.time.milliTimestamp();
    manager.stop();
    try std.testing.expect(std.time.milliTimestamp() - started < 500);
}

test "concurrent local channels pause at shared transport capacity and resume intact" {
    const alloc = std.testing.allocator;
    const port = try testUnusedPort();
    const manager = try Manager.init(alloc, .{ .sock = "/not-dialled" }, &.{.{ .local_port = port, .remote_port = 80 }});
    defer manager.stop();

    const transport_pair = try testSocketPair();
    defer std.posix.close(transport_pair[1]);
    var tr: client.Transport = .{ .link = .{ .fd = transport_pair[0] } };
    defer tr.close();
    var wire = try Wire.init(alloc, &tr);
    defer wire.deinit();
    const kernel_fill = try fillSocket(transport_pair[0]);

    const frame_len = proto.frame_header_len + proto.forward_id_len + proto.forward_data_max;
    const prefix_len = transport_queue_max - control_reserve - frame_len;
    try wire.output.resize(alloc, prefix_len);
    @memset(wire.output.items, 0x5a);

    const channel_count = 6;
    var channel_peers: [channel_count]std.posix.fd_t = undefined;
    var made: usize = 0;
    defer for (channel_peers[0..made]) |fd| std.posix.close(fd);
    var payloads: [channel_count][proto.forward_data_max]u8 = undefined;
    for (0..channel_count) |i| {
        const pair = try testSocketPair();
        channel_peers[i] = pair[1];
        made += 1;
        @memset(&payloads[i], @intCast(i + 1));
        try proto.writeAllFd(pair[1], &payloads[i]);
        manager.channels[i] = .{ .id = @intCast(i + 1), .fd = pair[0], .opening = false, .send_credit = proto.forward_initial_credit };
    }

    // Model one poll snapshot that reported every channel readable.
    for (0..channel_count) |i| try std.testing.expect(manager.readLocal(&wire, i));
    try std.testing.expectEqual(prefix_len + frame_len, wire.pendingBytes());
    try std.testing.expect(wire.pendingBytes() <= transport_queue_max - control_reserve);
    for (0..channel_count) |i| try std.testing.expect(manager.channels[i] != null);

    try drainQueuedPrefix(&wire, transport_pair[1], kernel_fill + prefix_len);
    try wire.flush();
    var first = (try proto.readFrame(alloc, transport_pair[1])) orelse return error.UnexpectedEof;
    defer first.deinit(alloc);
    try std.testing.expectEqual(proto.MsgType.forward_data, first.type);
    try std.testing.expectEqual(@as(u32, 1), try proto.decodeForwardId(first.payload));
    try std.testing.expectEqualSlices(u8, &payloads[0], first.payload[proto.forward_id_len..]);

    for (1..channel_count) |i| {
        try std.testing.expect(manager.readLocal(&wire, i));
        try std.testing.expect(wire.pendingBytes() <= transport_queue_max - control_reserve);
        try wire.flush();
        var frame = (try proto.readFrame(alloc, transport_pair[1])) orelse return error.UnexpectedEof;
        defer frame.deinit(alloc);
        try std.testing.expectEqual(proto.MsgType.forward_data, frame.type);
        try std.testing.expectEqual(@as(u32, @intCast(i + 1)), try proto.decodeForwardId(frame.payload));
        try std.testing.expectEqualSlices(u8, &payloads[i], frame.payload[proto.forward_id_len..]);
    }
}