better RST?

This commit is contained in:
Stefan Wasilewski 2026-08-01 04:29:50 +04:00
parent aa66d1c5f4
commit 34cb6339a8
3 changed files with 276 additions and 27 deletions

View file

@ -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,

114
main.go
View file

@ -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()
}

View file

@ -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)
})
}