package ssh import ( "crypto/rand" "crypto/rsa" "net" "os" "path/filepath" "strings" "testing" cryptossh "golang.org/x/crypto/ssh" "golang.org/x/crypto/ssh/knownhosts" ) func TestAcceptNewHostKeyCallbackPersistsUnknownHost(t *testing.T) { key := testPublicKey(t) knownHosts := filepath.Join(t.TempDir(), "known_hosts") callback, err := acceptNewHostKeyCallback(Options{KnownHosts: knownHosts}) if err != nil { t.Fatalf("acceptNewHostKeyCallback() error = %v", err) } if err := callback("example.com:22", &net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 22}, key); err != nil { t.Fatalf("callback() error = %v", err) } data, err := os.ReadFile(knownHosts) if err != nil { t.Fatalf("read known_hosts: %v", err) } if !strings.Contains(string(data), "example.com") { t.Fatalf("known_hosts = %q, want example.com entry", data) } if err := callback("example.com:22", &net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 22}, key); err != nil { t.Fatalf("second callback() error = %v", err) } } func TestAcceptNewHostKeyCallbackReadOnlyDoesNotPersistUnknownHost(t *testing.T) { key := testPublicKey(t) knownHosts := filepath.Join(t.TempDir(), "known_hosts") callback, err := acceptNewHostKeyCallback(Options{ KnownHosts: knownHosts, ReadOnlyKnownHosts: true, }) if err != nil { t.Fatalf("acceptNewHostKeyCallback() error = %v", err) } if err := callback("example.com:22", &net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 22}, key); err != nil { t.Fatalf("callback() error = %v", err) } if _, err := os.Stat(knownHosts); !os.IsNotExist(err) { t.Fatalf("known_hosts stat error = %v, want not exist", err) } } func TestAcceptNewHostKeyCallbackRejectsChangedHostKey(t *testing.T) { first := testPublicKey(t) second := testPublicKey(t) knownHosts := filepath.Join(t.TempDir(), "known_hosts") if err := os.WriteFile(knownHosts, []byte(knownhosts.Line([]string{knownhosts.Normalize("example.com:22")}, first)+"\n"), 0o600); err != nil { t.Fatalf("write known_hosts: %v", err) } callback, err := acceptNewHostKeyCallback(Options{KnownHosts: knownHosts}) if err != nil { t.Fatalf("acceptNewHostKeyCallback() error = %v", err) } err = callback("example.com:22", &net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 22}, second) if err == nil || !strings.Contains(err.Error(), "has changed") { t.Fatalf("callback() error = %v, want changed host key", err) } } func TestAcceptNewHostKeyCallbackReadOnlyRejectsChangedHostKey(t *testing.T) { first := testPublicKey(t) second := testPublicKey(t) knownHosts := filepath.Join(t.TempDir(), "known_hosts") if err := os.WriteFile(knownHosts, []byte(knownhosts.Line([]string{knownhosts.Normalize("example.com:22")}, first)+"\n"), 0o600); err != nil { t.Fatalf("write known_hosts: %v", err) } callback, err := acceptNewHostKeyCallback(Options{ KnownHosts: knownHosts, ReadOnlyKnownHosts: true, }) if err != nil { t.Fatalf("acceptNewHostKeyCallback() error = %v", err) } err = callback("example.com:22", &net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 22}, second) if err == nil || !strings.Contains(err.Error(), "has changed") { t.Fatalf("callback() error = %v, want changed host key", err) } } func TestAcceptNewHostKeyCallbackRequiresWritableKnownHostsForUnknownHost(t *testing.T) { callback, err := acceptNewHostKeyCallback(Options{}) if err != nil { t.Fatalf("acceptNewHostKeyCallback() error = %v", err) } err = callback("example.com:22", &net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 22}, testPublicKey(t)) if err == nil || !strings.Contains(err.Error(), "no writable known_hosts path") { t.Fatalf("callback() error = %v, want no writable known_hosts path", err) } } func TestAcceptNewHostKeyCallbackReadOnlyAllowsMissingKnownHosts(t *testing.T) { callback, err := acceptNewHostKeyCallback(Options{ReadOnlyKnownHosts: true}) if err != nil { t.Fatalf("acceptNewHostKeyCallback() error = %v", err) } if err := callback("example.com:22", &net.TCPAddr{IP: net.ParseIP("192.0.2.10"), Port: 22}, testPublicKey(t)); err != nil { t.Fatalf("callback() error = %v", err) } } func TestStrictHostKeyCallbackRequiresKnownHosts(t *testing.T) { _, err := hostKeyCallback(Options{HostKeyPolicy: HostKeyPolicyStrict}) if err == nil || !strings.Contains(err.Error(), "known_hosts is required") { t.Fatalf("hostKeyCallback() error = %v, want known_hosts required", err) } } func testPublicKey(t *testing.T) cryptossh.PublicKey { t.Helper() privateKey, err := rsa.GenerateKey(rand.Reader, 2048) if err != nil { t.Fatalf("generate key: %v", err) } publicKey, err := cryptossh.NewPublicKey(&privateKey.PublicKey) if err != nil { t.Fatalf("new public key: %v", err) } return publicKey }