From 34cb6339a80c31f2b023d1474636c67bdd0c2b80 Mon Sep 17 00:00:00 2001 From: Stefan Wasilewski Date: Sat, 1 Aug 2026 04:29:50 +0400 Subject: [PATCH] better RST? --- README.md | 13 ++++ main.go | 114 ++++++++++++++++++++++++++++----- main_test.go | 176 +++++++++++++++++++++++++++++++++++++++++++++++---- 3 files changed, 276 insertions(+), 27 deletions(-) diff --git a/README.md b/README.md index 2484729..56d3fbb 100644 --- a/README.md +++ b/README.md @@ -98,6 +98,19 @@ plist, then remove it once the state directory has been populated. So an unavailable target means `connection refused` on the local port, not a connect that immediately EOFs, and clients back off the way they would against a genuinely down service. +- **One failed connection closes the port.** A client that reconnects the + instant its connection breaks would otherwise beat the next probe and be + accepted into a forward with nothing behind it. So any connection that fails + against the target unbinds the port immediately, and it stays unbound until a + probe says the target is back. That probe is scheduled straight away, so a + one-off failure against a healthy target costs a few milliseconds of refusal, + not a whole interval. Other live connections are left alone — one failure + stops new work but isn't enough to declare the target dead. +- **Dials are bounded by `--probe-timeout`.** A tailnet peer that is routable + but dead neither accepts nor refuses, so an unbounded dial parks forever + holding a local socket that nothing can reclaim: the probe loop only tears + down live connections when a probe *fails*, and a recovering target makes the + probe succeed. Forwarded connections use the same timeout as probes. - **Target loss drops live connections.** A tailnet peer can vanish without the userspace TCP stack ever erroring on an established connection, which leaves local sockets hanging and apps waiting on a dead link. When a probe fails, diff --git a/main.go b/main.go index 64753b2..0895b22 100644 --- a/main.go +++ b/main.go @@ -248,7 +248,7 @@ func (p *proxy) markDown(cause error) { p.mu.Unlock() if was != stateDown { - log.Printf("tsproxy: %s -> %s: target unreachable: %v (refusing connections, dropped %d in flight)", + log.Printf("tsproxy: %s -> %s: target unreachable: %v (refusing connections, reset %d in flight)", p.local, p.target, cause, len(conns)) } if ln != nil { @@ -260,6 +260,34 @@ func (p *proxy) markDown(cause error) { } } +// suspend unbinds the local port after a single connection to the target +// failed, and leaves it unbound until a probe says the target is back. A +// client that reconnects the instant its connection breaks -- which is what +// clients do -- would otherwise be accepted into a forward with nothing behind +// it, so it must be refused rather than let in. +// +// Unlike markDown this does not touch other live connections: one failure is +// enough to stop admitting new work, but not enough to declare the target dead +// and reset connections that are still healthy. The probe it schedules decides +// that, and either rebinds the port or tears everything down. +func (p *proxy) suspend() { + p.mu.Lock() + ln := p.ln + p.ln = nil + was := p.state + p.state = stateDown + p.mu.Unlock() + + if ln != nil { + ln.Close() + } + if was == stateUp { + log.Printf("tsproxy: %s -> %s: connection to target failed, refusing connections pending probe", + p.local, p.target) + } + p.nudge() +} + func (p *proxy) accept(ctx context.Context, ln net.Listener) { for { c, err := ln.Accept() @@ -324,30 +352,84 @@ func abort(c net.Conn) { c.Close() } +// targetConn records whether the target end of a forwarded connection failed, +// so teardown can tell a broken target from a client that merely went away. +// Only the former should take the local port down. +// +// It embeds the net.Conn interface rather than a concrete type on purpose: that +// keeps ReadFrom/WriteTo off the method set, so io.Copy cannot take a fast path +// that bypasses these wrappers. +type targetConn struct { + net.Conn + + mu sync.Mutex + err error +} + +func (t *targetConn) note(err error) { + if err == nil || err == io.EOF { + return // a clean EOF is the target finishing normally, not failing + } + t.mu.Lock() + defer t.mu.Unlock() + if t.err == nil { + t.err = err + } +} + +func (t *targetConn) Read(b []byte) (int, error) { + n, err := t.Conn.Read(b) + t.note(err) + return n, err +} + +func (t *targetConn) Write(b []byte) (int, error) { + n, err := t.Conn.Write(b) + t.note(err) + return n, err +} + +func (t *targetConn) failure() error { + t.mu.Lock() + defer t.mu.Unlock() + return t.err +} + func (p *proxy) handle(ctx context.Context, in net.Conn) { defer p.untrack(in) - out, err := p.dial.Dial(ctx, "tcp", p.target) + // Bound the dial. A tailnet peer that is routable but dead accepts nothing + // and refuses nothing, so an unbounded dial parks here forever, holding a + // local socket open with no way out: the probe loop only tears down live + // connections when a probe fails, and a target that recovers makes the + // probe succeed. The connection would hang for good. + dialCtx, cancel := context.WithTimeout(ctx, p.timeout) + c, err := p.dial.Dial(dialCtx, "tcp", p.target) + cancel() // governs the dial only; the returned conn outlives it if err != nil { - log.Printf("dial %s: %v", p.target, err) - p.nudge() + log.Printf("tsproxy: %s -> %s: dial: %v", p.local, p.target, err) abort(in) + p.suspend() return } - defer out.Close() + out := &targetConn{Conn: c} + defer c.Close() - // Which direction ended, and whether it ended cleanly. A clean EOF from the - // target is passed through as a FIN -- that is exactly what the client - // would have seen connecting directly. An error is passed through as an - // RST, for the same reason. - type result struct{ err error } - done := make(chan result, 2) - go func() { _, err := io.Copy(out, in); done <- result{err} }() - go func() { _, err := io.Copy(in, out); done <- result{err} }() + // A clean EOF from the target is forwarded as a FIN -- exactly what the + // client would have seen connecting directly. Anything else is forwarded + // as an RST, and if the target was at fault, also takes the port down. + done := make(chan error, 2) + go func() { _, err := io.Copy(out, in); done <- err }() + go func() { _, err := io.Copy(in, out); done <- err }() - if r := <-done; r.err != nil { + copyErr := <-done + targetErr := out.failure() + if copyErr != nil || targetErr != nil { abort(in) - return + } else { + in.Close() + } + if targetErr != nil { + p.suspend() } - in.Close() } diff --git a/main_test.go b/main_test.go index 4c23f42..a7f93ea 100644 --- a/main_test.go +++ b/main_test.go @@ -11,16 +11,25 @@ import ( "time" ) +// how a fakeTarget responds to a dial. +type dialMode int + +const ( + modeUp dialMode = iota // connect to the echo server + modeDown // fail immediately + modeHang // block until the caller gives up, like a routable but dead peer +) + // fakeTarget stands in for a tailnet target. Its echo server always runs; the -// reachable flag only gates *new* dials. That models the failure this is meant -// to catch: a tailnet peer disappears, new dials fail, but connections already -// established just hang forever with no error from the userspace TCP stack. +// mode only gates *new* dials. That models the failure this is meant to catch: +// a tailnet peer disappears, new dials fail or hang, but connections already +// established just sit there with no error from the userspace TCP stack. type fakeTarget struct { ln net.Listener - mu sync.Mutex - reachable bool - accepted []net.Conn + mu sync.Mutex + mode dialMode + accepted []net.Conn } func newFakeTarget(t *testing.T) *fakeTarget { @@ -29,7 +38,7 @@ func newFakeTarget(t *testing.T) *fakeTarget { if err != nil { t.Fatalf("fake target listen: %v", err) } - ft := &fakeTarget{ln: ln, reachable: true} + ft := &fakeTarget{ln: ln, mode: modeUp} t.Cleanup(func() { ln.Close() }) go func() { for { @@ -61,18 +70,30 @@ func (f *fakeTarget) killAccepted(reset bool) { } } -func (f *fakeTarget) setReachable(v bool) { +func (f *fakeTarget) setMode(m dialMode) { f.mu.Lock() defer f.mu.Unlock() - f.reachable = v + f.mode = m +} + +func (f *fakeTarget) setReachable(v bool) { + if v { + f.setMode(modeUp) + } else { + f.setMode(modeDown) + } } func (f *fakeTarget) Dial(ctx context.Context, network, addr string) (net.Conn, error) { f.mu.Lock() - ok := f.reachable + mode := f.mode f.mu.Unlock() - if !ok { + switch mode { + case modeDown: return nil, errors.New("no route to host") + case modeHang: + <-ctx.Done() + return nil, ctx.Err() } var d net.Dialer return d.DialContext(ctx, network, f.ln.Addr().String()) @@ -301,3 +322,136 @@ func TestTargetHardCloseIsNoticedImmediately(t *testing.T) { }) } } + +// startProxyWith runs a proxy with explicit timings, for tests that need the +// probe loop held back so a pass can only come from the connection path. +func startProxyWith(t *testing.T, ft *fakeTarget, local string, interval, timeout time.Duration) *proxy { + t.Helper() + p := &proxy{ + dial: ft, + local: local, + target: "target:1234", + interval: interval, + timeout: timeout, + recheck: make(chan struct{}, 1), + conns: make(map[net.Conn]struct{}), + } + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { defer close(done); p.run(ctx) }() + t.Cleanup(func() { cancel(); <-done }) + return p +} + +// A target that is routable but dead makes the dial block rather than fail. +// The client must not be left parked on a socket whose handler is stuck in +// Dial -- nothing else would ever reclaim it, because the probe loop only +// tears down live connections when a probe FAILS, and a target that comes back +// makes the probe succeed. +func TestHangingDialDoesNotStrandClient(t *testing.T) { + ft := newFakeTarget(t) + local := freePort(t) + // Probe interval far beyond the test, so the probe loop cannot be what + // rescues the connection. + startProxyWith(t, ft, local, 30*time.Second, 300*time.Millisecond) + waitFor(t, "local port to accept", func() bool { return !localPortRefused(local) }) + + ft.setMode(modeHang) + + c, err := net.Dial("tcp", local) + if err != nil { + t.Fatalf("dial local: %v", err) + } + defer c.Close() + + start := time.Now() + c.SetReadDeadline(time.Now().Add(5 * time.Second)) + buf := make([]byte, 1) + _, err = c.Read(buf) + elapsed := time.Since(start) + if err == nil { + t.Fatal("read succeeded against a hanging target") + } + if isTimeout(err) { + t.Fatalf("client stranded on a hanging dial (blocked %s); dial must be bounded", elapsed) + } + t.Logf("client released after %s (err: %v)", elapsed, err) + if elapsed > 2*time.Second { + t.Errorf("client held for %s; want release near the probe timeout", elapsed) + } + // Recovery is deliberately not asserted here: the 30s interval that makes + // this test meaningful is also long enough that rebinding would be a race + // against the probe loop's next tick. TestUnreachableTargetRefusesConnections + // and TestPortReturnsAfterTransientFailure cover it. +} + +// What the user hit: a client reconnects the instant its connection breaks, +// beating the probe. It must be refused, not accepted into a dead forward. +func TestFailedConnectionClosesPortBeforeProbe(t *testing.T) { + ft := newFakeTarget(t) + local := freePort(t) + // Long interval and a long probe timeout: once the target starts hanging, + // a probe cannot complete within the assertion window, so the port closing + // can only be the connection path doing it. + startProxyWith(t, ft, local, 30*time.Second, 5*time.Second) + waitFor(t, "local port to accept", func() bool { return !localPortRefused(local) }) + + c, err := net.Dial("tcp", local) + if err != nil { + t.Fatalf("dial local: %v", err) + } + defer c.Close() + if _, err := c.Write([]byte("ping")); err != nil { + t.Fatalf("write: %v", err) + } + buf := make([]byte, 4) + c.SetReadDeadline(time.Now().Add(3 * time.Second)) + if _, err := io.ReadFull(c, buf); err != nil { + t.Fatalf("read echo: %v", err) + } + + ft.setMode(modeHang) + start := time.Now() + ft.killAccepted(true) // target drops the connection, as on process death + + deadline := time.Now().Add(2 * time.Second) + for !localPortRefused(local) { + if time.Now().After(deadline) { + t.Fatal("local port still accepting after the target dropped a connection") + } + time.Sleep(5 * time.Millisecond) + } + t.Logf("port closed %s after target dropped the connection", time.Since(start)) +} + +// Suspending on one failure must not strand the port: if the target is in fact +// fine, the probe that suspension schedules has to bring it straight back. +func TestPortReturnsAfterTransientFailure(t *testing.T) { + ft := newFakeTarget(t) + local := freePort(t) + // Long interval: recovery must come from the probe suspension schedules, + // not from the next scheduled tick. + startProxyWith(t, ft, local, 30*time.Second, time.Second) + waitFor(t, "local port to accept", func() bool { return !localPortRefused(local) }) + + c, err := net.Dial("tcp", local) + if err != nil { + t.Fatalf("dial local: %v", err) + } + if _, err := c.Write([]byte("ping")); err != nil { + t.Fatalf("write: %v", err) + } + buf := make([]byte, 4) + c.SetReadDeadline(time.Now().Add(3 * time.Second)) + if _, err := io.ReadFull(c, buf); err != nil { + t.Fatalf("read echo: %v", err) + } + + // The target stays healthy; this one connection just breaks. + ft.killAccepted(true) + c.Close() + + waitFor(t, "local port to come back without waiting for a tick", func() bool { + return !localPortRefused(local) + }) +}