88 lines
2.8 KiB
Go
88 lines
2.8 KiB
Go
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 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 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 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
|
|
}
|