diff --git a/README.md b/README.md index 20cb1a4..2484729 100644 --- a/README.md +++ b/README.md @@ -103,6 +103,16 @@ plist, then remove it once the state directory has been populated. local sockets hanging and apps waiting on a dead link. When a probe fails, every connection on that forward is closed so clients see the drop and reconnect. Detection takes up to `--probe-interval` + `--probe-timeout`. +- **Failures are reset, not closed.** How a connection ends is forwarded + faithfully. A target that closes cleanly gives the local client a FIN (an + ordinary EOF); a target that resets, errors, or disappears gives it an RST. + This matters for clients that hold idle connections: a FIN leaves the socket + writable, so a pooled client's *next* write still succeeds and only the one + after it fails, whereas an RST fails the very next read or write. It also + keeps a truncated response from looking like a complete one. The trade-off + is that an RST discards whatever was still queued in the send buffer — a + stream torn down this way was already incomplete, so flagging it beats + delivering a partial result that looks whole. - **Probes are real connections.** Each probe opens and immediately closes a TCP connection to the target. Chatty services may log these; raise `--probe-interval` to quiet them down, at the cost of slower detection. diff --git a/main.go b/main.go index 839a9e3..64753b2 100644 --- a/main.go +++ b/main.go @@ -254,8 +254,9 @@ func (p *proxy) markDown(cause error) { if ln != nil { ln.Close() } + // RST, not FIN: these connections did not end, they broke. See abort. for _, c := range conns { - c.Close() + abort(c) } } @@ -307,20 +308,46 @@ func (p *proxy) untrack(c net.Conn) { delete(p.conns, c) } +// abort closes c with a TCP RST instead of a FIN. A FIN says "the peer +// finished normally", which is a lie when the target broke: an idle client +// won't notice until its next write, and a client reading a response with no +// declared length cannot tell a truncated body from a complete one. An RST +// fails the peer's next read or write immediately and unambiguously. +// +// This discards anything still queued in the send buffer, which is the point: +// a stream that ended this way was incomplete regardless, and flagging it +// beats delivering a partial result that looks whole. +func abort(c net.Conn) { + if tc, ok := c.(*net.TCPConn); ok { + tc.SetLinger(0) + } + c.Close() +} + func (p *proxy) handle(ctx context.Context, in net.Conn) { defer p.untrack(in) - defer in.Close() out, err := p.dial.Dial(ctx, "tcp", p.target) if err != nil { log.Printf("dial %s: %v", p.target, err) p.nudge() + abort(in) return } defer out.Close() - done := make(chan struct{}, 2) - go func() { io.Copy(out, in); done <- struct{}{} }() - go func() { io.Copy(in, out); done <- struct{}{} }() - <-done + // 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} }() + + if r := <-done; r.err != nil { + abort(in) + return + } + in.Close() } diff --git a/main_test.go b/main_test.go index c43a961..4c23f42 100644 --- a/main_test.go +++ b/main_test.go @@ -6,6 +6,7 @@ import ( "io" "net" "sync" + "syscall" "testing" "time" ) @@ -19,6 +20,7 @@ type fakeTarget struct { mu sync.Mutex reachable bool + accepted []net.Conn } func newFakeTarget(t *testing.T) *fakeTarget { @@ -35,12 +37,30 @@ func newFakeTarget(t *testing.T) *fakeTarget { if err != nil { return } + ft.mu.Lock() + ft.accepted = append(ft.accepted, c) + ft.mu.Unlock() go func() { io.Copy(c, c); c.Close() }() } }() return ft } +// killAccepted hard-closes every connection the target has accepted, as if the +// service process were killed. reset picks RST (SO_LINGER 0) over a clean FIN. +func (f *fakeTarget) killAccepted(reset bool) { + f.mu.Lock() + conns := f.accepted + f.accepted = nil + f.mu.Unlock() + for _, c := range conns { + if reset { + c.(*net.TCPConn).SetLinger(0) + } + c.Close() + } +} + func (f *fakeTarget) setReachable(v bool) { f.mu.Lock() defer f.mu.Unlock() @@ -188,14 +208,18 @@ func TestTargetLossClosesLocalConnection(t *testing.T) { // nothing but the probe loop can notice. ft.setReachable(false) - // The client's blocking read must return, and reasonably promptly. + // The client's blocking read must return, promptly, and as a reset -- a + // plain EOF here would tell the client the stream ended normally. c.SetReadDeadline(time.Now().Add(3 * time.Second)) - if _, err := c.Read(buf); err == nil { + _, err = c.Read(buf) + if err == nil { t.Fatal("read succeeded after target loss; want the connection closed") - } else if errors.Is(err, io.EOF) || isTimeout(err) { - if isTimeout(err) { - t.Fatal("read blocked after target loss; local socket was never closed") - } + } + if isTimeout(err) { + t.Fatal("read blocked after target loss; local socket was never closed") + } + if !isReset(err) { + t.Errorf("read err = %v, want a connection reset (a clean EOF would falsely signal a complete stream)", err) } // New connects must be refused too. @@ -206,3 +230,74 @@ func isTimeout(err error) bool { var ne net.Error return errors.As(err, &ne) && ne.Timeout() } + +func isReset(err error) bool { + return errors.Is(err, syscall.ECONNRESET) +} + +// A target process that dies drops its connections with a FIN or an RST. The +// local client must see that within a round trip -- it must NOT have to wait +// for the reachability probe, which is the slow fallback for a peer that +// vanishes without any TCP signal at all. +func TestTargetHardCloseIsNoticedImmediately(t *testing.T) { + for _, tc := range []struct { + name string + reset bool + }{ + {"fin", false}, + {"rst", true}, + } { + t.Run(tc.name, func(t *testing.T) { + ft := newFakeTarget(t) + local := freePort(t) + // Probe interval far longer than the assertion window, so a pass + // can only mean the per-connection path noticed. + p := &proxy{ + dial: ft, + local: local, + target: "target:1234", + interval: 30 * time.Second, + timeout: time.Second, + 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) }() + defer func() { cancel(); <-done }() + + 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) + } + + start := time.Now() + ft.killAccepted(tc.reset) + + c.SetReadDeadline(time.Now().Add(5 * time.Second)) + _, err = c.Read(buf) + elapsed := time.Since(start) + if err == nil { + t.Fatal("read succeeded after target hard-close") + } + if isTimeout(err) { + t.Fatalf("client never noticed target hard-close (blocked %s)", elapsed) + } + t.Logf("client noticed after %s (err: %v)", elapsed, err) + if elapsed > time.Second { + t.Errorf("client took %s to notice; want sub-second", elapsed) + } + }) + } +}