| [ef1c1c7] | 1 | package db
|
|---|
| 2 |
|
|---|
| 3 | import (
|
|---|
| 4 | "fmt"
|
|---|
| 5 | "net"
|
|---|
| 6 | "os"
|
|---|
| 7 | "time"
|
|---|
| 8 |
|
|---|
| 9 | "golang.org/x/crypto/ssh"
|
|---|
| 10 | )
|
|---|
| 11 |
|
|---|
| 12 | // sshDialer implements pq.Dialer by opening every database connection
|
|---|
| 13 | // through an SSH client, so no separate `ssh -L` tunnel is needed.
|
|---|
| 14 | type sshDialer struct {
|
|---|
| 15 | client *ssh.Client
|
|---|
| 16 | }
|
|---|
| 17 |
|
|---|
| 18 | func (d sshDialer) Dial(network, address string) (net.Conn, error) {
|
|---|
| 19 | return d.client.Dial(network, address)
|
|---|
| 20 | }
|
|---|
| 21 |
|
|---|
| 22 | func (d sshDialer) DialTimeout(network, address string, _ time.Duration) (net.Conn, error) {
|
|---|
| 23 | return d.client.Dial(network, address)
|
|---|
| 24 | }
|
|---|
| 25 |
|
|---|
| 26 | // newSSHDialer connects to the SSH server described by the SSH_* variables.
|
|---|
| 27 | // Authentication uses SSH_KEY (path to a private key) and/or SSH_PASSWORD.
|
|---|
| 28 | func newSSHDialer(host string) (sshDialer, error) {
|
|---|
| 29 | port := getenv("SSH_PORT", "22")
|
|---|
| 30 | user := os.Getenv("SSH_USER")
|
|---|
| 31 | if user == "" {
|
|---|
| 32 | return sshDialer{}, fmt.Errorf("SSH_HOST is set but SSH_USER is empty")
|
|---|
| 33 | }
|
|---|
| 34 |
|
|---|
| 35 | var auth []ssh.AuthMethod
|
|---|
| 36 | if keyPath := os.Getenv("SSH_KEY"); keyPath != "" {
|
|---|
| 37 | key, err := os.ReadFile(keyPath)
|
|---|
| 38 | if err != nil {
|
|---|
| 39 | return sshDialer{}, fmt.Errorf("read SSH_KEY: %w", err)
|
|---|
| 40 | }
|
|---|
| 41 | var signer ssh.Signer
|
|---|
| 42 | if phrase := os.Getenv("SSH_KEY_PASSPHRASE"); phrase != "" {
|
|---|
| 43 | signer, err = ssh.ParsePrivateKeyWithPassphrase(key, []byte(phrase))
|
|---|
| 44 | } else {
|
|---|
| 45 | signer, err = ssh.ParsePrivateKey(key)
|
|---|
| 46 | }
|
|---|
| 47 | if err != nil {
|
|---|
| 48 | return sshDialer{}, fmt.Errorf("parse SSH_KEY: %w", err)
|
|---|
| 49 | }
|
|---|
| 50 | auth = append(auth, ssh.PublicKeys(signer))
|
|---|
| 51 | }
|
|---|
| 52 | if pass := os.Getenv("SSH_PASSWORD"); pass != "" {
|
|---|
| 53 | auth = append(auth, ssh.Password(pass))
|
|---|
| 54 | }
|
|---|
| 55 | if len(auth) == 0 {
|
|---|
| 56 | return sshDialer{}, fmt.Errorf("SSH_HOST is set but neither SSH_KEY nor SSH_PASSWORD is")
|
|---|
| 57 | }
|
|---|
| 58 |
|
|---|
| 59 | config := &ssh.ClientConfig{
|
|---|
| 60 | User: user,
|
|---|
| 61 | Auth: auth,
|
|---|
| 62 | // Host key is not pinned; acceptable for a course prototype that only
|
|---|
| 63 | // talks to the faculty server.
|
|---|
| 64 | HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
|---|
| 65 | Timeout: 10 * time.Second,
|
|---|
| 66 | }
|
|---|
| 67 | client, err := ssh.Dial("tcp", net.JoinHostPort(host, port), config)
|
|---|
| 68 | if err != nil {
|
|---|
| 69 | return sshDialer{}, fmt.Errorf("ssh dial %s@%s:%s: %w", user, host, port, err)
|
|---|
| 70 | }
|
|---|
| 71 | return sshDialer{client: client}, nil
|
|---|
| 72 | }
|
|---|