diff --git a/README.md b/README.md index a522bea..20cb1a4 100644 --- a/README.md +++ b/README.md @@ -59,6 +59,8 @@ go build -o /usr/local/bin/tsproxy . | `--name` | `-n` | `tsproxy` | Hostname advertised on the tailnet. | | `--dir` | | `~/.config/tsproxy/` | State directory (node identity, keys). | | `--verbose` | `-v` | `false` | Verbose tsnet logging. | +| `--probe-interval` | | `5s` | How often each target is probed for reachability. | +| `--probe-timeout` | | `3s` | How long a probe may take before the target counts as unreachable. | Target can be any MagicDNS name, short hostname, or tailnet IP. @@ -88,8 +90,22 @@ plist, then remove it once the state directory has been populated. - **One node, many forwards.** All `--forward` rules share a single tsnet identity, so you get one device in the admin console and one ACL subject. -- **Startup is strict.** If any local listener fails to bind, the process - exits — partial success is confusing under a supervisor. +- **Startup is strict.** Every local address is bound once at startup as a + check; if any fails, the process exits — partial success is confusing under + a supervisor. +- **The local port tracks the target.** Each forward dials its target every + `--probe-interval` and only keeps the local port bound while that succeeds. + 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. +- **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, + every connection on that forward is closed so clients see the drop and + reconnect. Detection takes up to `--probe-interval` + `--probe-timeout`. +- **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. - **State directory matters.** Losing `~/.config/tsproxy//` means the node re-registers on next launch and will need a fresh `TS_AUTHKEY`. diff --git a/main.go b/main.go index 63c0324..839a9e3 100644 --- a/main.go +++ b/main.go @@ -15,6 +15,7 @@ import ( "path/filepath" "strings" "sync" + "time" flag "github.com/spf13/pflag" "tailscale.com/tsnet" @@ -50,6 +51,10 @@ func main() { hostname = flag.StringP("name", "n", "tsproxy", "hostname this node advertises on the tailnet") dir = flag.String("dir", "", "state directory (default: ~/.config/tsproxy/)") verbose = flag.BoolP("verbose", "v", false, "verbose tsnet logging") + interval = flag.Duration("probe-interval", 5*time.Second, + "how often to probe each target for reachability") + timeout = flag.Duration("probe-timeout", 3*time.Second, + "how long a target probe may take before the target counts as unreachable") ) flag.Parse() @@ -60,6 +65,12 @@ func main() { if err != nil { log.Fatal(err) } + if *interval <= 0 { + log.Fatal("--probe-interval must be positive") + } + if *timeout <= 0 { + log.Fatal("--probe-timeout must be positive") + } stateDir := *dir if stateDir == "" { @@ -73,6 +84,18 @@ func main() { log.Fatalf("mkdir state: %v", err) } + // Bind every local address up front so a typo or a port clash still kills + // the process at startup rather than a supervisor restart later. The + // listeners are handed straight back; from here on each proxy binds and + // unbinds its own port to track target reachability. + for _, f := range fwds { + ln, err := net.Listen("tcp", f.local) + if err != nil { + log.Fatalf("listen %s: %v", f.local, err) + } + ln.Close() + } + srv := &tsnet.Server{ Hostname: *hostname, Dir: stateDir, @@ -89,36 +112,209 @@ func main() { var wg sync.WaitGroup for _, f := range fwds { - ln, err := net.Listen("tcp", f.local) - if err != nil { - log.Fatalf("listen %s: %v", f.local, err) + p := &proxy{ + dial: srv, + local: f.local, + target: f.target, + interval: *interval, + timeout: *timeout, + recheck: make(chan struct{}, 1), + conns: make(map[net.Conn]struct{}), } log.Printf("tsproxy: %s -> %s (via tailnet as %q)", f.local, f.target, *hostname) wg.Add(1) - go func(ln net.Listener, target string) { + go func() { defer wg.Done() - serve(ctx, srv, ln, target) - }(ln, f.target) + p.run(ctx) + }() } wg.Wait() } -func serve(ctx context.Context, srv *tsnet.Server, ln net.Listener, target string) { +// dialer is the subset of *tsnet.Server that a proxy needs. +type dialer interface { + Dial(ctx context.Context, network, addr string) (net.Conn, error) +} + +// target reachability, as last observed by the probe loop. +type state int + +const ( + stateUnknown state = iota + stateUp + stateDown +) + +// proxy forwards one local address to one tailnet target, and owns the local +// listener: the port is only bound while the target is known to be reachable, +// so clients get a connection refusal (not a connect-then-EOF) while it isn't. +type proxy struct { + dial dialer + local string + target string + interval time.Duration + timeout time.Duration + recheck chan struct{} // nudges the probe loop to re-probe immediately + + mu sync.Mutex + state state + ln net.Listener // nil while the target is down + conns map[net.Conn]struct{} +} + +func (p *proxy) run(ctx context.Context) { + t := time.NewTicker(p.interval) + defer t.Stop() for { - c, err := ln.Accept() - if err != nil { - log.Printf("accept %s: %v", ln.Addr(), err) - return + if err := p.probe(ctx); err != nil { + p.markDown(err) + } else { + p.markUp(ctx) + } + select { + case <-ctx.Done(): + p.markDown(ctx.Err()) + return + case <-t.C: + case <-p.recheck: } - go handle(ctx, srv, c, target) } } -func handle(ctx context.Context, srv *tsnet.Server, in net.Conn, target string) { - defer in.Close() - out, err := srv.Dial(ctx, "tcp", target) +// probe dials the target and hangs up. It is the only evidence we have that +// the target is alive: a tailnet peer can disappear without the userspace TCP +// stack ever reporting an error on an established connection. +func (p *proxy) probe(ctx context.Context) error { + ctx, cancel := context.WithTimeout(ctx, p.timeout) + defer cancel() + c, err := p.dial.Dial(ctx, "tcp", p.target) if err != nil { - log.Printf("dial %s: %v", target, err) + return err + } + return c.Close() +} + +// nudge asks the probe loop to re-probe now rather than at the next tick. +func (p *proxy) nudge() { + select { + case p.recheck <- struct{}{}: + default: + } +} + +// markUp binds the local port if it isn't bound already. +func (p *proxy) markUp(ctx context.Context) { + p.mu.Lock() + defer p.mu.Unlock() + + was := p.state + p.state = stateUp + if p.ln != nil { + return + } + ln, err := net.Listen("tcp", p.local) + if err != nil { + // Someone else holds the port. Stay down and retry on the next tick. + p.state = stateDown + if was != stateDown { + log.Printf("tsproxy: %s -> %s: listen: %v (retrying every %s)", p.local, p.target, err, p.interval) + } + return + } + p.ln = ln + if was == stateUp { + // Listener was lost without the target going down; already logged. + log.Printf("tsproxy: %s -> %s: listening again", p.local, p.target) + } else { + log.Printf("tsproxy: %s -> %s: target reachable, accepting connections", p.local, p.target) + } + go p.accept(ctx, ln) +} + +// markDown unbinds the local port so further connects are refused, and closes +// every connection already in flight so clients see the drop and reconnect +// instead of blocking forever on a socket whose far end is gone. +func (p *proxy) markDown(cause error) { + p.mu.Lock() + was := p.state + p.state = stateDown + ln := p.ln + p.ln = nil + conns := make([]net.Conn, 0, len(p.conns)) + for c := range p.conns { + conns = append(conns, c) + } + clear(p.conns) + p.mu.Unlock() + + if was != stateDown { + log.Printf("tsproxy: %s -> %s: target unreachable: %v (refusing connections, dropped %d in flight)", + p.local, p.target, cause, len(conns)) + } + if ln != nil { + ln.Close() + } + for _, c := range conns { + c.Close() + } +} + +func (p *proxy) accept(ctx context.Context, ln net.Listener) { + for { + c, err := ln.Accept() + if err != nil { + // Either markDown closed this listener (expected) or it failed on + // its own. Drop it if it's still the live one; the probe loop will + // rebind on the next tick. + p.mu.Lock() + current := p.ln == ln + if current { + p.ln = nil + } + p.mu.Unlock() + if current { + log.Printf("tsproxy: %s -> %s: accept: %v", p.local, p.target, err) + ln.Close() + p.nudge() + } + return + } + if !p.track(c) { + // Raced with markDown: the target went away between Accept and + // here, so this connection is already condemned. + c.Close() + continue + } + go p.handle(ctx, c) + } +} + +// track registers a connection for mass close on target loss. It reports false +// if the target is already down, in which case the connection is not tracked. +func (p *proxy) track(c net.Conn) bool { + p.mu.Lock() + defer p.mu.Unlock() + if p.state != stateUp { + return false + } + p.conns[c] = struct{}{} + return true +} + +func (p *proxy) untrack(c net.Conn) { + p.mu.Lock() + defer p.mu.Unlock() + delete(p.conns, c) +} + +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() return } defer out.Close() diff --git a/main_test.go b/main_test.go new file mode 100644 index 0000000..c43a961 --- /dev/null +++ b/main_test.go @@ -0,0 +1,208 @@ +package main + +import ( + "context" + "errors" + "io" + "net" + "sync" + "testing" + "time" +) + +// 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. +type fakeTarget struct { + ln net.Listener + + mu sync.Mutex + reachable bool +} + +func newFakeTarget(t *testing.T) *fakeTarget { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("fake target listen: %v", err) + } + ft := &fakeTarget{ln: ln, reachable: true} + t.Cleanup(func() { ln.Close() }) + go func() { + for { + c, err := ln.Accept() + if err != nil { + return + } + go func() { io.Copy(c, c); c.Close() }() + } + }() + return ft +} + +func (f *fakeTarget) setReachable(v bool) { + f.mu.Lock() + defer f.mu.Unlock() + f.reachable = v +} + +func (f *fakeTarget) Dial(ctx context.Context, network, addr string) (net.Conn, error) { + f.mu.Lock() + ok := f.reachable + f.mu.Unlock() + if !ok { + return nil, errors.New("no route to host") + } + var d net.Dialer + return d.DialContext(ctx, network, f.ln.Addr().String()) +} + +// freePort returns an address that is currently bindable. +func freePort(t *testing.T) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("free port: %v", err) + } + addr := ln.Addr().String() + ln.Close() + return addr +} + +func startProxy(t *testing.T, ft *fakeTarget, local string) *proxy { + t.Helper() + p := &proxy{ + dial: ft, + local: local, + target: "target:1234", + interval: 20 * time.Millisecond, + 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) }() + t.Cleanup(func() { + cancel() + <-done + }) + return p +} + +// waitFor polls cond until it holds or the deadline passes. +func waitFor(t *testing.T, what string, cond func() bool) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + if cond() { + return + } + time.Sleep(5 * time.Millisecond) + } + t.Fatalf("timed out waiting for %s", what) +} + +func localPortRefused(local string) bool { + c, err := net.DialTimeout("tcp", local, 500*time.Millisecond) + if err != nil { + return true + } + c.Close() + return false +} + +// An unreachable target must mean a refused connect, not a successful connect +// followed by an immediate EOF. +func TestUnreachableTargetRefusesConnections(t *testing.T) { + ft := newFakeTarget(t) + ft.setReachable(false) + local := freePort(t) + startProxy(t, ft, local) + + // Give the probe loop several ticks; the port must never come up. + time.Sleep(200 * time.Millisecond) + if !localPortRefused(local) { + t.Fatal("connect succeeded while target was unreachable; want refused") + } + + // And it must start accepting once the target comes back. + ft.setReachable(true) + waitFor(t, "local port to accept after target recovery", func() bool { + return !localPortRefused(local) + }) +} + +func TestForwardsData(t *testing.T) { + ft := newFakeTarget(t) + local := freePort(t) + startProxy(t, ft, local) + 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) + } + if string(buf) != "ping" { + t.Fatalf("echo = %q, want %q", buf, "ping") + } +} + +// The reported bug: when the proxy loses the target, an already-established +// local socket must be closed so the client's read fails and it reconnects, +// rather than blocking forever on a half-dead connection. +func TestTargetLossClosesLocalConnection(t *testing.T) { + ft := newFakeTarget(t) + local := freePort(t) + startProxy(t, ft, local) + 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() + + // Establish that the connection is live and forwarding. + 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) + } + + // Target vanishes. The echo server still holds the far end open, so + // nothing but the probe loop can notice. + ft.setReachable(false) + + // The client's blocking read must return, and reasonably promptly. + c.SetReadDeadline(time.Now().Add(3 * time.Second)) + if _, err := c.Read(buf); 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") + } + } + + // New connects must be refused too. + waitFor(t, "local port to be refused", func() bool { return localPortRefused(local) }) +} + +func isTimeout(err error) bool { + var ne net.Error + return errors.As(err, &ne) && ne.Timeout() +}