Something went wrong. Try again.
Monorepo for Tangled tangled.org
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193package mill
import ( "context" "crypto/ed25519" "crypto/rand" "encoding/pem" "fmt" "net" "os" "path/filepath" "sync" "time"
"github.com/gliderlabs/ssh" gossh "golang.org/x/crypto/ssh" "tangled.org/core/spindle/observability")
const ( jumpIdleTimeout = 5 * time.Minute jumpMaxTimeout = 24 * time.Hour maxJumpConnectionsPerIP = 8)
type jumpContextKey string
const jumpRouteOpened jumpContextKey = "route-opened"
func (m *Mill) ServeJump(ctx context.Context, listenAddr, hostKeyPath string, executorPort uint32, maxConnections int) { if listenAddr == "" { return } srv, err := m.newJumpServer(hostKeyPath, executorPort, maxConnections) if err != nil { m.l.Error("setup debug ssh jump server", "err", err) return }
go func() { <-ctx.Done() _ = srv.Close() }()
m.l.Info("starting debug ssh jump server", "address", listenAddr) srv.Addr = listenAddr if err := srv.ListenAndServe(); err != nil && err != ssh.ErrServerClosed { m.l.Error("debug ssh jump server stopped", "err", err) }}
func (m *Mill) newJumpServer(hostKeyPath string, executorPort uint32, maxConnections int) (*ssh.Server, error) { if executorPort == 0 { executorPort = 2223 } if maxConnections <= 0 { return nil, fmt.Errorf("max jump connections must be greater than zero") } if err := ensureJumpHostKey(hostKeyPath); err != nil { return nil, fmt.Errorf("prepare jump host key: %w", err) } limiter := newJumpConnectionLimiter(maxConnections, maxJumpConnectionsPerIP) limiter.metrics = m.metrics m.metrics.RecordJumpLimit(int64(maxConnections)) srv := &ssh.Server{ PublicKeyHandler: func(ctx ssh.Context, _ ssh.PublicKey) bool { allowed := ctx.User() == "debug" if !allowed { m.metrics.RecordJumpRejection("unauthorized_user") } return allowed }, ConnCallback: func(ctx ssh.Context, conn net.Conn) net.Conn { if !limiter.acquire(conn.RemoteAddr()) { _ = conn.Close() return conn } go func() { <-ctx.Done() limiter.release(conn.RemoteAddr()) }() return conn }, LocalPortForwardingCallback: func(ctx ssh.Context, host string, port uint32) bool { if port != executorPort { m.metrics.RecordJumpRejection("invalid_port") return false } if !m.hasLiveExecutor(host) { m.metrics.RecordJumpRejection("executor_unavailable") return false } ctx.Lock() defer ctx.Unlock() if opened, _ := ctx.Value(jumpRouteOpened).(bool); opened { m.metrics.RecordJumpRejection("route_already_open") return false } ctx.SetValue(jumpRouteOpened, true) return true }, ChannelHandlers: map[string]ssh.ChannelHandler{ "direct-tcpip": ssh.DirectTCPIPHandler, }, IdleTimeout: jumpIdleTimeout, MaxTimeout: jumpMaxTimeout, } if err := srv.SetOption(ssh.HostKeyFile(hostKeyPath)); err != nil { return nil, fmt.Errorf("load jump host key: %w", err) } return srv, nil}
type jumpConnectionLimiter struct { mu sync.Mutex total int perIP map[string]int maxTotal int maxPerIP int metrics *observability.Metrics}
func newJumpConnectionLimiter(maxTotal, maxPerIP int) *jumpConnectionLimiter { return &jumpConnectionLimiter{ perIP: make(map[string]int), maxTotal: maxTotal, maxPerIP: maxPerIP, }}
func (l *jumpConnectionLimiter) acquire(addr net.Addr) bool { host := jumpRemoteHost(addr) l.mu.Lock() defer l.mu.Unlock() metrics := l.metrics if l.total >= l.maxTotal { metrics.RecordJumpRejection("max_total_reached") return false } if l.perIP[host] >= l.maxPerIP { metrics.RecordJumpRejection("max_per_ip_reached") return false } l.total++ l.perIP[host]++ metrics.RecordJumpActive(int64(l.total)) return true}
func (l *jumpConnectionLimiter) release(addr net.Addr) { host := jumpRemoteHost(addr) l.mu.Lock() defer l.mu.Unlock() l.total-- l.perIP[host]-- if l.perIP[host] == 0 { delete(l.perIP, host) } l.metrics.RecordJumpActive(int64(l.total))}
func jumpRemoteHost(addr net.Addr) string { if addr == nil { return "" } host, _, err := net.SplitHostPort(addr.String()) if err != nil { return addr.String() } return host}
func ensureJumpHostKey(path string) error { if _, err := os.Stat(path); err == nil { return nil } else if !os.IsNotExist(err) { return fmt.Errorf("stat host key: %w", err) }
_, privateKey, err := ed25519.GenerateKey(rand.Reader) if err != nil { return fmt.Errorf("generate host key: %w", err) } block, err := gossh.MarshalPrivateKey(privateKey, "") if err != nil { return fmt.Errorf("marshal host key: %w", err) } if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { return fmt.Errorf("create host key directory: %w", err) } return os.WriteFile(path, pem.EncodeToMemory(block), 0o600)}