Files
distributor/internal/adapters/ssh/hostkeys_test.go

139 lines
4.6 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 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
}