// tsproxy forwards local TCP ports to host:port targets on your tailnet, // joining the tailnet itself via tsnet (no system tailscaled required). // // First run: set TS_AUTHKEY (from https://login.tailscale.com/admin/settings/keys) // to register the node. State is persisted under --dir, so subsequent runs // don't need the auth key. package main import ( "context" "fmt" "io" "log" "net" "os" "path/filepath" "strings" "sync" "sync/atomic" "time" flag "github.com/spf13/pflag" "tailscale.com/tsnet" ) type forward struct { local string target string } func parseForwards(specs []string) ([]forward, error) { out := make([]forward, 0, len(specs)) for _, s := range specs { l, t, ok := strings.Cut(s, "=") if !ok || l == "" || t == "" { return nil, &parseError{spec: s} } out = append(out, forward{local: l, target: t}) } return out, nil } type parseError struct{ spec string } func (e *parseError) Error() string { return "bad --forward spec " + e.spec + " (want LOCAL=TARGET, e.g. 127.0.0.1:8080=myhost:80)" } func main() { var ( forwards = flag.StringArrayP("forward", "f", nil, "forward rule LOCAL=TARGET (repeatable), e.g. --forward 127.0.0.1:8080=myhost:80") 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 re-probe a target that is down (a target that is up is never probed)") timeout = flag.Duration("probe-timeout", 10*time.Second, "how long a reachability probe may take before the target counts as down") dialTimeout = flag.Duration("dial-timeout", 30*time.Second, "how long a forwarded connection may take to reach the target before it is failed") trace = flag.Bool("trace", false, "log every accept, dial, probe and teardown (for diagnosing stalls)") idle = flag.Duration("idle-timeout", 5*time.Minute, "close a forwarded connection after this long with no traffic either way (0 disables)") ) flag.Parse() if len(*forwards) == 0 { log.Fatal("at least one --forward LOCAL=TARGET is required") } fwds, err := parseForwards(*forwards) 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") } if *dialTimeout <= 0 { log.Fatal("--dial-timeout must be positive") } if *idle < 0 { log.Fatal("--idle-timeout must not be negative (0 disables)") } stateDir := *dir if stateDir == "" { home, err := os.UserHomeDir() if err != nil { log.Fatalf("home dir: %v", err) } stateDir = filepath.Join(home, ".config", "tsproxy", *hostname) } if err := os.MkdirAll(stateDir, 0o700); err != nil { 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, } if !*verbose { srv.Logf = func(string, ...any) {} } defer srv.Close() ctx := context.Background() if _, err := srv.Up(ctx); err != nil { log.Fatalf("tsnet up: %v", err) } var wg sync.WaitGroup for _, f := range fwds { 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{}), trace: *trace, idle: *idle, dialWait: *dialTimeout, } log.Printf("tsproxy: %s -> %s (via tailnet as %q)", f.local, f.target, *hostname) wg.Add(1) go func() { defer wg.Done() p.run(ctx) }() } wg.Wait() } // 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 // budget for a reachability probe dialWait time.Duration // budget for a forwarded connection to reach the target idle time.Duration // tear down a connection after this long with no traffic; 0 disables recheck chan struct{} // nudges the probe loop to re-probe immediately trace bool // log every accept, dial, probe and teardown seq atomic.Uint64 // connection counter, for correlating log lines mu sync.Mutex state state ln net.Listener // nil while the target is down conns map[net.Conn]struct{} } func (p *proxy) logf(format string, args ...any) { log.Printf("tsproxy: %s -> %s: %s", p.local, p.target, fmt.Sprintf(format, args...)) } // tracef logs only under --trace: the per-connection and per-probe detail you // want while diagnosing a stall, and not otherwise. func (p *proxy) tracef(format string, args ...any) { if p.trace { p.logf(format, args...) } } func (p *proxy) liveConns() int { p.mu.Lock() defer p.mu.Unlock() return len(p.conns) } // run probes the target and binds or unbinds the local port to match. // // Probing happens only while the target is believed down: once it is up, real // connections are the health signal and further probes would be pure noise // against the target. One probe runs at startup to decide whether to bind at // all, and after that the loop sits idle until something reports a failure. func (p *proxy) run(ctx context.Context) { for { start := time.Now() err := p.probe(ctx) took := time.Since(start).Round(time.Millisecond) if err != nil { p.tracef("probe failed after %s: %v (%d live)", took, err, p.liveConns()) p.markDown(err) } else { p.tracef("probe ok in %s (%d live)", took, p.liveConns()) p.markUp(ctx) } // Retry on a timer only while down. While up, wait to be woken by a // connection that failed -- there is nothing to poll for. var retry <-chan time.Time var timer *time.Timer if !p.isUp() { timer = time.NewTimer(p.interval) retry = timer.C } select { case <-ctx.Done(): if timer != nil { timer.Stop() } p.markDown(ctx.Err()) return case <-retry: case <-p.recheck: if timer != nil { timer.Stop() } } } } func (p *proxy) isUp() bool { p.mu.Lock() defer p.mu.Unlock() return p.state == stateUp } // probe dials the target and hangs up. It is how a down target is found to be // back: 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 { 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. p.logf("PORT OPEN: listening again on %s", p.local) } else { p.logf("PORT OPEN: target reachable, listening on %s", p.local) } 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 { p.logf("PORT CLOSED: target unreachable: %v (refusing connections, resetting %d in flight)", cause, len(conns)) } else { p.tracef("still down: %v (%d in flight to reset)", cause, len(conns)) } if ln != nil { ln.Close() } // RST, not FIN: these connections did not end, they broke. See abort. for _, c := range conns { abort(c) } } // 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() // Log before closing: closing wakes the accept loop, whose own trace line // would otherwise print first and read as if it caused this. if was == stateUp { p.logf("PORT CLOSED: a connection to the target failed, refusing connections pending probe") } else { p.tracef("suspend: port was already closed") } if ln != nil { ln.Close() } p.nudge() } 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 { // Nobody is accepting on a bound port now -- that would hang a // client in connect(), so say so loudly. p.logf("PORT CLOSED: accept loop died on %s: %v", p.local, err) ln.Close() p.nudge() } else { p.tracef("accept loop exited on %s (listener already replaced)", p.local) } return } id := p.seq.Add(1) p.tracef("[conn %d] accepted from %s", id, c.RemoteAddr()) if !p.track(c) { // Raced with the port closing: the target went down between Accept // and here, so this connection is already condemned. Reset rather // than close, or the client sees a successful connect followed by a // clean EOF -- the exact thing the port closing exists to prevent. p.tracef("[conn %d] reset: target went down during accept", id) abort(c) continue } go p.handle(ctx, c, id) } } // 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) } // 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() } // 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 } // A dial slower than this is reported without --trace. It is not a fraction of // --dial-timeout on purpose: the budget is about when to give up, this is about // how long a client will sit connected but unserved, which is a much smaller // number and is what actually breaks applications. const slowDial = time.Second // which side of a forwarded connection finished first, and what it moved. type direction struct { name string bytes int64 err error } // countingReader records that bytes moved, for the idle watchdog. Both // directions share one counter: traffic either way means the connection is // alive, so an active download must not let the quiet upload side time out. type countingReader struct { r io.Reader moved *atomic.Int64 } func (c *countingReader) Read(b []byte) (int, error) { n, err := c.r.Read(b) if n > 0 { c.moved.Add(int64(n)) } return n, err } // watchIdle tears a connection down once no bytes have moved either way for // p.idle. With probing suppressed while the target is up, this is what catches // a target that accepted a connection and then went silent -- otherwise the // client waits forever on a socket whose far end is gone. // // It sets idled before closing so the teardown can tell this apart from a // target failure: a connection nobody was using is no evidence the target is // down, and must not take the local port with it. func (p *proxy) watchIdle(id uint64, in, out net.Conn, moved *atomic.Int64, idled *atomic.Bool, stop <-chan struct{}) { tick := p.idle / 4 if tick <= 0 { tick = p.idle } t := time.NewTicker(tick) defer t.Stop() last := moved.Load() lastChange := time.Now() for { select { case <-stop: return case <-t.C: if n := moved.Load(); n != last { last, lastChange = n, time.Now() continue } if quiet := time.Since(lastChange); quiet >= p.idle { idled.Store(true) p.logf("[conn %d] no traffic for %s, closing", id, quiet.Round(time.Second)) abort(in) out.Close() return } } } } func (p *proxy) handle(ctx context.Context, in net.Conn, id uint64) { defer p.untrack(in) opened := time.Now() // 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. // This budget is deliberately not --probe-timeout. A probe is a health // check that costs nothing to retry, so it can be impatient; a forwarded // dial has a client waiting on it and would rather wait than fail. Tailnet // dials are also wildly variable -- tens of milliseconds on a warm path, // seconds when the path has to be renegotiated or fall back to DERP -- so // a probe-sized budget here fails connections to a perfectly healthy // target. p.tracef("[conn %d] dialing %s (timeout %s)", id, p.target, p.dialWait) dialStart := time.Now() dialCtx, cancel := context.WithTimeout(ctx, p.dialWait) c, err := p.dial.Dial(dialCtx, "tcp", p.target) cancel() // governs the dial only; the returned conn outlives it dialTook := time.Since(dialStart).Round(time.Millisecond) if err != nil { p.logf("[conn %d] dial failed after %s: %v -- resetting client", id, dialTook, err) abort(in) p.suspend() return } if dialTook > slowDial { // Worth seeing without --trace. This is measured against what a client // will sit through, not against the dial budget: the client is already // connected -- the local handshake completed the moment it called // connect() -- so every second spent here is a second it waits with no // idea anything is wrong, and long enough will trip its own timeout. p.logf("[conn %d] SLOW DIAL: connected in %s -- the client has been waiting that long", id, dialTook) } else { p.tracef("[conn %d] connected to target in %s", id, dialTook) } out := &targetConn{Conn: c} defer c.Close() var moved atomic.Int64 var idled atomic.Bool if p.idle > 0 { stop := make(chan struct{}) defer close(stop) go p.watchIdle(id, in, out, &moved, &idled, stop) } // 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 direction, 2) go func() { n, err := io.Copy(out, &countingReader{in, &moved}) done <- direction{"client->target", n, err} }() go func() { n, err := io.Copy(in, &countingReader{out, &moved}) done <- direction{"target->client", n, err} }() first := <-done targetErr := out.failure() lived := time.Since(opened).Round(time.Millisecond) switch { case idled.Load(): // The watchdog already reset the client. Deliberately no suspend: an // unused connection says nothing about whether the target is healthy. p.tracef("[conn %d] idle-closed after %s, %d bytes total", id, lived, moved.Load()) case targetErr != nil: p.logf("[conn %d] target failed after %s: %v (%s moved %d bytes) -- resetting client", id, lived, targetErr, first.name, first.bytes) abort(in) p.suspend() case first.err != nil: // The target is wired through targetConn, so any failure of its own // would have set targetErr and taken the branch above. Reaching here // means the local client is what broke -- it gave up, or went away. p.logf("[conn %d] client gave up after %s: %v (%s moved %d bytes)", id, lived, first.err, first.name, first.bytes) abort(in) default: p.tracef("[conn %d] closed cleanly after %s: %s ended, %d bytes", id, lived, first.name, first.bytes) in.Close() } }