326 lines
8 KiB
Go
326 lines
8 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"
|
|
"io"
|
|
"log"
|
|
"net"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync"
|
|
"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 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()
|
|
|
|
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")
|
|
}
|
|
|
|
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{}),
|
|
}
|
|
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
|
|
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 {
|
|
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:
|
|
}
|
|
}
|
|
}
|
|
|
|
// 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 {
|
|
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()
|
|
|
|
done := make(chan struct{}, 2)
|
|
go func() { io.Copy(out, in); done <- struct{}{} }()
|
|
go func() { io.Copy(in, out); done <- struct{}{} }()
|
|
<-done
|
|
}
|