tsproxy/main.go
Stefan Wasilewski 1e59753238 more rst fixes 2
2026-08-01 05:37:46 +04:00

634 lines
19 KiB
Go

// 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/<name>)")
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()
}
}