130 lines
3.1 KiB
Go
130 lines
3.1 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"
|
|
|
|
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")
|
|
)
|
|
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)
|
|
}
|
|
|
|
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)
|
|
}
|
|
|
|
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 {
|
|
ln, err := net.Listen("tcp", f.local)
|
|
if err != nil {
|
|
log.Fatalf("listen %s: %v", f.local, err)
|
|
}
|
|
log.Printf("tsproxy: %s -> %s (via tailnet as %q)", f.local, f.target, *hostname)
|
|
wg.Add(1)
|
|
go func(ln net.Listener, target string) {
|
|
defer wg.Done()
|
|
serve(ctx, srv, ln, target)
|
|
}(ln, f.target)
|
|
}
|
|
wg.Wait()
|
|
}
|
|
|
|
func serve(ctx context.Context, srv *tsnet.Server, ln net.Listener, target string) {
|
|
for {
|
|
c, err := ln.Accept()
|
|
if err != nil {
|
|
log.Printf("accept %s: %v", ln.Addr(), err)
|
|
return
|
|
}
|
|
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)
|
|
if err != nil {
|
|
log.Printf("dial %s: %v", target, err)
|
|
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
|
|
}
|