package ssh import ( "context" "fmt" "net" "os" "os/signal" "syscall" "time" "golang.org/x/term" "git.tukangketik.id/swanadiva/hostkeeper/internal/errors" "git.tukangketik.id/swanadiva/hostkeeper/internal/models" cryptossh "golang.org/x/crypto/ssh" ) // Client represents an SSH client type Client struct { host *models.Host timeout time.Duration client *cryptossh.Client config *cryptossh.ClientConfig } // NewClient creates a new SSH client func NewClient(host *models.Host, timeout time.Duration) *Client { return &Client{ host: host, timeout: timeout, } } // Connect establishes an SSH connection func (c *Client) Connect(ctx context.Context) error { // Create SSH configuration if err := c.setupConfig(); err != nil { return fmt.Errorf("failed to setup SSH config: %w", err) } // Create connection context with timeout connCtx, cancel := context.WithTimeout(ctx, c.timeout) defer cancel() // Establish TCP connection address := fmt.Sprintf("%s:%d", c.host.Hostname, c.host.Port) conn, err := c.dialTCP(connCtx, address) if err != nil { return errors.HandleSSHError(err) } // Establish SSH connection over TCP sshConn, chans, reqs, err := cryptossh.NewClientConn(conn, address, c.config) if err != nil { conn.Close() return errors.HandleSSHError(err) } c.client = cryptossh.NewClient(sshConn, chans, reqs) return nil } // dialTCP establishes a TCP connection func (c *Client) dialTCP(ctx context.Context, address string) (net.Conn, error) { d := net.Dialer{} return d.DialContext(ctx, "tcp", address) } // setupConfig creates SSH client configuration func (c *Client) setupConfig() error { config := &cryptossh.ClientConfig{ User: c.host.Username, HostKeyCallback: cryptossh.InsecureIgnoreHostKey(), //nolint:gosec // Phase 1 - will be improved in Phase 2 Timeout: c.timeout, } // Configure authentication methods authMethods, err := c.getAuthMethods() if err != nil { return fmt.Errorf("failed to setup authentication: %w", err) } config.Auth = authMethods c.config = config return nil } // Execute runs a command on the remote server func (c *Client) Execute(_ context.Context, cmd string) (string, error) { if c.client == nil { return "", fmt.Errorf("not connected to server") } session, err := c.client.NewSession() if err != nil { return "", fmt.Errorf("failed to create session: %w", err) } defer session.Close() output, err := session.CombinedOutput(cmd) if err != nil { return string(output), fmt.Errorf("command execution failed: %w", err) } return string(output), nil } // Shell opens an interactive shell session func (c *Client) Shell() error { if c.client == nil { return fmt.Errorf("not connected to server") } session, err := c.client.NewSession() if err != nil { return fmt.Errorf("failed to create session: %w", err) } defer session.Close() // Get current terminal state fd := int(os.Stdin.Fd()) oldState, err := term.MakeRaw(fd) if err != nil { return fmt.Errorf("failed to set raw terminal: %w", err) } defer term.Restore(fd, oldState) // Set up terminal modes modes := cryptossh.TerminalModes{ cryptossh.ECHO: 1, cryptossh.TTY_OP_ISPEED: 14400, cryptossh.TTY_OP_OSPEED: 14400, } // Get terminal size width, height, err := term.GetSize(fd) if err != nil { width = 80 height = 24 } // Request PTY if err := session.RequestPty("xterm-256color", height, width, modes); err != nil { return fmt.Errorf("failed to request PTY: %w", err) } // Handle window changes sigwinch := make(chan os.Signal, 1) signal.Notify(sigwinch, os.Signal(syscall.SIGWINCH)) go func() { for range sigwinch { w, h, err := term.GetSize(fd) if err != nil { continue } session.WindowChange(h, w) } }() defer signal.Stop(sigwinch) // Link I/O session.Stdin = os.Stdin session.Stdout = os.Stdout session.Stderr = os.Stderr // Start shell if err := session.Shell(); err != nil { return fmt.Errorf("failed to start shell: %w", err) } // Wait for shell to exit if err := session.Wait(); err != nil { if exitErr, ok := err.(*cryptossh.ExitError); ok { if exitErr.ExitStatus() != 0 { return fmt.Errorf("shell exited with status %d", exitErr.ExitStatus()) } return nil } return fmt.Errorf("shell session error: %w", err) } return nil } // Close closes the SSH connection func (c *Client) Close() error { if c.client != nil { return c.client.Close() } return nil } // GetClient returns the underlying SSH client func (c *Client) GetClient() *cryptossh.Client { return c.client } // IsConnected returns true if the client has an active connection func (c *Client) IsConnected() bool { return c.client != nil }