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..]);
}
}