diff --git a/toolbox/cmd/root.go b/toolbox/cmd/root.go index 121ce6f..040c3dc 100644 --- a/toolbox/cmd/root.go +++ b/toolbox/cmd/root.go @@ -2,13 +2,36 @@ package cmd import ( "os" + "path/filepath" + "github.com/charmbracelet/log" "github.com/spf13/cobra" ) +var ( + hostsFile string + host string + sshUser string + sshKey string + sshKnownHosts string +) + +func init() { + log.SetReportTimestamp(false) + + rootCmd.PersistentFlags().StringVar(&hostsFile, "hosts-file", "", "Path to hosts.json file") + rootCmd.PersistentFlags().StringVar(&host, "host", "", "Host name to connect to (e.g., kube-1)") + rootCmd.PersistentFlags().StringVar(&sshUser, "ssh-user", "root", "SSH user") + rootCmd.PersistentFlags().StringVar(&sshKey, "ssh-key", defaultSSHKey(), "Path to SSH private key") + rootCmd.PersistentFlags().StringVar(&sshKnownHosts, "ssh-known-hosts", defaultKnownHostsFile(), "Path to SSH known_hosts file") +} + var rootCmd = &cobra.Command{ Use: "toolbox", Short: "CLI tools for managing cloudlab infrastructure", + PersistentPreRunE: func(cmd *cobra.Command, args []string) error { + return nil + }, } func Execute() { @@ -16,3 +39,19 @@ func Execute() { os.Exit(1) } } + +func defaultSSHKey() string { + home, err := os.UserHomeDir() + if err != nil { + return "" + } + return filepath.Join(home, ".ssh", "id_ed25519") +} + +func defaultKnownHostsFile() string { + home, err := os.UserHomeDir() + if err != nil { + return "" + } + return filepath.Join(home, ".ssh", "known_hosts") +} diff --git a/toolbox/go.mod b/toolbox/go.mod index 10935eb..591b69b 100644 --- a/toolbox/go.mod +++ b/toolbox/go.mod @@ -3,6 +3,30 @@ module github.com/khuedoan/cloudlab/toolbox go 1.25.5 require ( + github.com/charmbracelet/log v0.4.2 + github.com/hashicorp/vault/api v1.22.0 github.com/spf13/cobra v1.10.2 golang.org/x/crypto v0.47.0 ) + +require ( + github.com/cenkalti/backoff/v4 v4.3.0 // indirect + github.com/go-jose/go-jose/v4 v4.1.1 // indirect + github.com/go-logfmt/logfmt v0.6.0 // indirect + github.com/hashicorp/errwrap v1.1.0 // indirect + github.com/hashicorp/go-cleanhttp v0.5.2 // indirect + github.com/hashicorp/go-multierror v1.1.1 // indirect + github.com/hashicorp/go-retryablehttp v0.7.8 // indirect + github.com/hashicorp/go-rootcerts v1.0.2 // indirect + github.com/hashicorp/go-secure-stdlib/parseutil v0.2.0 // indirect + github.com/hashicorp/go-secure-stdlib/strutil v0.1.2 // indirect + github.com/hashicorp/go-sockaddr v1.0.7 // indirect + github.com/hashicorp/hcl v1.0.1-vault-7 // indirect + github.com/mitchellh/go-homedir v1.1.0 // indirect + github.com/mitchellh/mapstructure v1.5.0 // indirect + github.com/ryanuber/go-glob v1.0.0 // indirect + golang.org/x/net v0.48.0 // indirect + golang.org/x/sys v0.40.0 // indirect + golang.org/x/text v0.33.0 // indirect + golang.org/x/time v0.12.0 // indirect +) diff --git a/toolbox/internal/cluster/client.go b/toolbox/internal/cluster/client.go new file mode 100644 index 0000000..1b893a4 --- /dev/null +++ b/toolbox/internal/cluster/client.go @@ -0,0 +1,135 @@ +package cluster + +import ( + "context" + "encoding/base64" + "fmt" + "strings" + "time" + + "github.com/hashicorp/vault/api" +) + +const ( + vaultNamespace = "vault" + vaultService = "svc/vault" + vaultPort = 8200 +) + +type ClientConfig struct { + HostsFile string + Host string + SSHUser string + SSHKey string + SSHKnownHosts string + Timeout time.Duration +} + +type Client struct { + conn *Connector + vault *api.Client +} + +func NewClient(ctx context.Context, cfg ClientConfig) (*Client, error) { + if cfg.Timeout == 0 { + cfg.Timeout = 30 * time.Second + } + + connectCtx, cancel := context.WithTimeout(ctx, cfg.Timeout) + defer cancel() + + hostAddr, err := LoadHost(cfg.HostsFile, cfg.Host) + if err != nil { + return nil, fmt.Errorf("load host: %w", err) + } + + conn, err := Connect(SSHConfig{ + Host: hostAddr, + User: cfg.SSHUser, + KeyPath: cfg.SSHKey, + KnownHostsPath: cfg.SSHKnownHosts, + Timeout: cfg.Timeout, + }) + if err != nil { + return nil, fmt.Errorf("connect: %w", err) + } + + token, err := getVaultToken(connectCtx, conn) + if err != nil { + conn.Close() + return nil, fmt.Errorf("get vault token: %w", err) + } + + vaultTunnel, err := conn.Forward(connectCtx, ServiceConfig{ + Namespace: vaultNamespace, + Name: vaultService, + Port: vaultPort, + }) + if err != nil { + conn.Close() + return nil, fmt.Errorf("forward vault: %w", err) + } + + vaultClient, err := newVaultClient(vaultTunnel.LocalAddr, token) + if err != nil { + conn.Close() + return nil, fmt.Errorf("create vault client: %w", err) + } + + return &Client{ + conn: conn, + vault: vaultClient, + }, nil +} + +func (c *Client) Vault() *api.Client { + return c.vault +} + +func (c *Client) Forward(ctx context.Context, svc ServiceConfig) (*ServiceTunnel, error) { + return c.conn.Forward(ctx, svc) +} + +func (c *Client) RunCommand(cmd string) ([]byte, error) { + return c.conn.RunCommand(cmd) +} + +func (c *Client) RunCommandContext(ctx context.Context, cmd string) ([]byte, error) { + return c.conn.RunCommandContext(ctx, cmd) +} + +func (c *Client) Close() error { + return c.conn.Close() +} + +func getVaultToken(ctx context.Context, conn *Connector) (string, error) { + cmd := fmt.Sprintf( + `kubectl get secret vault-unseal-keys -n %s -o template='{{ index .data "vault-root" }}'`, + vaultNamespace, + ) + + output, err := conn.RunCommandContext(ctx, cmd) + if err != nil { + return "", err + } + + token, err := base64.StdEncoding.DecodeString(strings.TrimSpace(string(output))) + if err != nil { + return "", fmt.Errorf("decode token: %w", err) + } + + return string(token), nil +} + +func newVaultClient(addr, token string) (*api.Client, error) { + config := api.DefaultConfig() + config.Address = "http://" + addr + + client, err := api.NewClient(config) + if err != nil { + return nil, fmt.Errorf("create client: %w", err) + } + client.SetToken(token) + + return client, nil +} diff --git a/toolbox/internal/cluster/cluster.go b/toolbox/internal/cluster/cluster.go index a4222e9..f401a09 100644 --- a/toolbox/internal/cluster/cluster.go +++ b/toolbox/internal/cluster/cluster.go @@ -3,9 +3,13 @@ package cluster import ( "context" "encoding/json" + "errors" "fmt" + "io" + "net" "os" "strings" + "sync" "time" "golang.org/x/crypto/ssh" @@ -13,8 +17,9 @@ import ( ) const ( - defaultSSHPort = 22 - defaultSSHTimeout = 10 * time.Second + defaultSSHPort = 22 + defaultSSHTimeout = 10 * time.Second + healthCheckInterval = 500 * time.Millisecond ) type SSHConfig struct { @@ -25,8 +30,29 @@ type SSHConfig struct { Timeout time.Duration } +type ServiceConfig struct { + Namespace string + Name string + Port int +} + type Connector struct { sshClient *ssh.Client + tunnels []*tunnel + mu sync.Mutex +} + +type ServiceTunnel struct { + LocalAddr string +} + +type tunnel struct { + listener net.Listener + session *ssh.Session + localPort int + remotePort int + done chan struct{} + closeOnce sync.Once } func Connect(cfg SSHConfig) (*Connector, error) { @@ -39,7 +65,64 @@ func Connect(cfg SSHConfig) (*Connector, error) { return nil, fmt.Errorf("ssh connect: %w", err) } - return &Connector{sshClient: sshClient}, nil + return &Connector{ + sshClient: sshClient, + tunnels: make([]*tunnel, 0), + }, nil +} + +func (c *Connector) Forward(ctx context.Context, svc ServiceConfig) (*ServiceTunnel, error) { + c.mu.Lock() + defer c.mu.Unlock() + + session, err := c.sshClient.NewSession() + if err != nil { + return nil, fmt.Errorf("create session: %w", err) + } + + cmd := fmt.Sprintf( + "exec kubectl port-forward %s -n %s %d:%d", + svc.Name, svc.Namespace, svc.Port, svc.Port, + ) + + if err := session.Start(cmd); err != nil { + session.Close() + return nil, fmt.Errorf("start port-forward: %w", err) + } + + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + session.Signal(ssh.SIGTERM) + session.Close() + return nil, fmt.Errorf("listen: %w", err) + } + + localPort := listener.Addr().(*net.TCPAddr).Port + done := make(chan struct{}) + + t := &tunnel{ + listener: listener, + session: session, + localPort: localPort, + remotePort: svc.Port, + done: done, + } + + go c.runTunnel(t) + + c.tunnels = append(c.tunnels, t) + + localAddr := fmt.Sprintf("127.0.0.1:%d", localPort) + if err := waitForTunnel(ctx, c.sshClient, localAddr, svc.Port); err != nil { + closeErr := c.closeTunnel(t) + c.tunnels = c.tunnels[:len(c.tunnels)-1] + if closeErr != nil { + return nil, fmt.Errorf("service not reachable: %w (cleanup: %v)", err, closeErr) + } + return nil, fmt.Errorf("service not reachable: %w", err) + } + + return &ServiceTunnel{LocalAddr: localAddr}, nil } func (c *Connector) RunCommand(cmd string) ([]byte, error) { @@ -77,14 +160,148 @@ func (c *Connector) RunCommandContext(ctx context.Context, cmd string) ([]byte, } func (c *Connector) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + + var errs []error + + for _, t := range c.tunnels { + if err := c.closeTunnel(t); err != nil { + errs = append(errs, err) + } + } + c.tunnels = nil + if c.sshClient != nil { if err := c.sshClient.Close(); err != nil { - return fmt.Errorf("close ssh: %w", err) + errs = append(errs, fmt.Errorf("close ssh: %w", err)) } } + + if len(errs) > 0 { + return errors.Join(errs...) + } return nil } +func (c *Connector) closeTunnel(t *tunnel) error { + var closeErr error + + t.closeOnce.Do(func() { + var errs []error + + close(t.done) + + if t.listener != nil { + if err := t.listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + errs = append(errs, fmt.Errorf("close listener: %w", err)) + } + t.listener = nil + } + + if t.session != nil { + _ = t.session.Signal(ssh.SIGTERM) + if err := t.session.Close(); err != nil && !errors.Is(err, io.EOF) { + errs = append(errs, fmt.Errorf("close session: %w", err)) + } + t.session = nil + } + + if len(errs) > 0 { + closeErr = errors.Join(errs...) + } + }) + + return closeErr +} + +func (c *Connector) runTunnel(t *tunnel) { + for { + select { + case <-t.done: + return + default: + } + + localConn, err := t.listener.Accept() + if err != nil { + if errors.Is(err, net.ErrClosed) { + return + } + continue + } + + go c.handleTunnelConn(localConn, t.remotePort) + } +} + +func (c *Connector) handleTunnelConn(localConn net.Conn, remotePort int) { + defer localConn.Close() + + remoteAddr := fmt.Sprintf("127.0.0.1:%d", remotePort) + remoteConn, err := c.sshClient.Dial("tcp", remoteAddr) + if err != nil { + return + } + defer remoteConn.Close() + + done := make(chan struct{}, 2) + + go func() { + io.Copy(remoteConn, localConn) + done <- struct{}{} + }() + + go func() { + io.Copy(localConn, remoteConn) + done <- struct{}{} + }() + + <-done +} + +func waitForTunnel(ctx context.Context, sshClient *ssh.Client, localAddr string, remotePort int) error { + dialer := &net.Dialer{Timeout: 2 * time.Second} + remoteAddr := fmt.Sprintf("127.0.0.1:%d", remotePort) + var lastErr error + + for { + select { + case <-ctx.Done(): + if lastErr != nil { + return fmt.Errorf("%w (last check: %v)", ctx.Err(), lastErr) + } + return ctx.Err() + default: + } + + localConn, err := dialer.DialContext(ctx, "tcp", localAddr) + if err != nil { + lastErr = fmt.Errorf("dial local %s: %w", localAddr, err) + } else { + localConn.Close() + + // Ensure the SSH-side endpoint is also reachable so we don't report + // readiness while the remote port-forward is still starting. + remoteConn, err := sshClient.Dial("tcp", remoteAddr) + if err == nil { + remoteConn.Close() + return nil + } + lastErr = fmt.Errorf("dial remote %s: %w", remoteAddr, err) + } + + select { + case <-ctx.Done(): + if lastErr != nil { + return fmt.Errorf("%w (last check: %v)", ctx.Err(), lastErr) + } + return ctx.Err() + case <-time.After(healthCheckInterval): + } + } +} + type HostInfo struct { IPv6Address string `json:"ipv6_address"` } @@ -143,6 +360,7 @@ func dialSSH(cfg SSHConfig) (*ssh.Client, error) { addr := fmt.Sprintf("%s:%d", cfg.Host, defaultSSHPort) if strings.Contains(cfg.Host, ":") { + // IPv6 addresses need brackets addr = fmt.Sprintf("[%s]:%d", cfg.Host, defaultSSHPort) }