package offsitekeys import ( "bytes" "context" "fmt" "net" "strconv" "time" "golang.org/x/crypto/ssh" ) // SSHDialer is the production Dialer: password (and keyboard-interactive) auth to the sub-account's // restricted shell, the host key checked against the provisioning-time fingerprint. The password is // never logged and never leaves this process. type SSHDialer struct { Timeout time.Duration // connect + per-command bound; 0 → 30 s } type sshShell struct { c *ssh.Client timeout time.Duration } func (d SSHDialer) Dial(ctx context.Context, t Target, password string) (Shell, error) { to := d.Timeout if to == 0 { to = 30 * time.Second } port := t.Port if port == 0 { port = 23 } want := t.Fingerprint cfg := &ssh.ClientConfig{ User: t.User, Auth: []ssh.AuthMethod{ ssh.Password(password), ssh.KeyboardInteractive(func(_, _ string, qs []string, _ []bool) ([]string, error) { ans := make([]string, len(qs)) for i := range ans { ans[i] = password } return ans, nil }), }, HostKeyCallback: func(_ string, _ net.Addr, key ssh.PublicKey) error { if got := ssh.FingerprintSHA256(key); got != want { return fmt.Errorf("offsitekeys: host key MISMATCH for %s (got %s, want %s)", t.Host, got, want) } return nil }, Timeout: to, } addr := net.JoinHostPort(t.Host, strconv.Itoa(port)) var nd net.Dialer dctx, cancel := context.WithTimeout(ctx, to) defer cancel() conn, err := nd.DialContext(dctx, "tcp", addr) if err != nil { return nil, fmt.Errorf("offsitekeys: dial %s: %w", addr, err) } _ = conn.SetDeadline(time.Now().Add(to)) cc, chans, reqs, err := ssh.NewClientConn(conn, addr, cfg) if err != nil { conn.Close() return nil, fmt.Errorf("offsitekeys: ssh handshake %s: %w", addr, err) } _ = conn.SetDeadline(time.Time{}) return &sshShell{c: ssh.NewClient(cc, chans, reqs), timeout: to}, nil } func (s *sshShell) Run(ctx context.Context, cmd string, stdin []byte) ([]byte, error) { sess, err := s.c.NewSession() if err != nil { return nil, err } defer sess.Close() if stdin != nil { sess.Stdin = bytes.NewReader(stdin) } var out, errb bytes.Buffer sess.Stdout, sess.Stderr = &out, &errb done := make(chan error, 1) go func() { done <- sess.Run(cmd) }() rctx, cancel := context.WithTimeout(ctx, s.timeout) defer cancel() select { case err := <-done: if err != nil { return out.Bytes(), fmt.Errorf("%q: %w: %s", cmd, err, bytes.TrimSpace(errb.Bytes())) } return out.Bytes(), nil case <-rctx.Done(): _ = sess.Close() return nil, fmt.Errorf("%q: %w", cmd, rctx.Err()) } } func (s *sshShell) Close() error { return s.c.Close() }