package main import ( "bytes" "errors" "io" "net" "os" "path/filepath" "strconv" "strings" "testing" "golang.org/x/crypto/ssh" "golang.org/x/crypto/ssh/knownhosts" ) // genClientKey makes a client key via keygen, writes it to a file for RSH_KEY, // and returns the private key path and the corresponding public key. func genClientKey(t *testing.T) (keyPath string, pub ssh.PublicKey) { t.Helper() var buf bytes.Buffer if err := keygen(&buf); err != nil { t.Fatal(err) } keyPath = filepath.Join(t.TempDir(), "id_ed25519") if err := os.WriteFile(keyPath, buf.Bytes(), 0o600); err != nil { t.Fatal(err) } signer, err := ssh.ParsePrivateKey(buf.Bytes()) if err != nil { t.Fatal(err) } return keyPath, signer.PublicKey() } func writeKnownHosts(t *testing.T, host string, port int, hostKey ssh.PublicKey) string { t.Helper() addr := knownhosts.Normalize(net.JoinHostPort(host, strconv.Itoa(port))) line := knownhosts.Line([]string{addr}, hostKey) p := filepath.Join(t.TempDir(), "known_hosts") if err := os.WriteFile(p, []byte(line+"\n"), 0o600); err != nil { t.Fatal(err) } return p } func setEnv(t *testing.T, key, kh string, port int) { t.Setenv("RSH_KEY", key) t.Setenv("RSH_KNOWN_HOSTS", kh) t.Setenv("RSH_PORT", strconv.Itoa(port)) } func TestTransportBridgesStdio(t *testing.T) { keyPath, pub := genClientKey(t) echo := func(_ string, stdin io.Reader, stdout, _ io.Writer) int { io.Copy(stdout, stdin) return 0 } srv := newTestServer(t, pub, echo) kh := writeKnownHosts(t, "127.0.0.1", srv.port(), srv.hostKey.PublicKey()) setEnv(t, keyPath, kh, srv.port()) want := []byte("the quick brown fox\x00\x01\x02 binary tail") var out bytes.Buffer if err := transport([]string{"-l", "u", "127.0.0.1", "cat"}, bytes.NewReader(want), &out, io.Discard); err != nil { t.Fatalf("transport: %v", err) } if !bytes.Equal(out.Bytes(), want) { t.Fatalf("bridged data mismatch:\n got %q\nwant %q", out.Bytes(), want) } } func TestTransportKeyFromData(t *testing.T) { keyPath, pub := genClientKey(t) echo := func(_ string, stdin io.Reader, stdout, _ io.Writer) int { io.Copy(stdout, stdin) return 0 } srv := newTestServer(t, pub, echo) kh := writeKnownHosts(t, "127.0.0.1", srv.port(), srv.hostKey.PublicKey()) keyData, err := os.ReadFile(keyPath) if err != nil { t.Fatal(err) } // RSH_KEY points nowhere; RSH_KEY_DATA must take precedence and work. t.Setenv("RSH_KEY", "/does/not/exist") t.Setenv("RSH_KEY_DATA", string(keyData)) t.Setenv("RSH_KNOWN_HOSTS", kh) t.Setenv("RSH_PORT", strconv.Itoa(srv.port())) want := []byte("via key data") var out bytes.Buffer if err := transport([]string{"-l", "u", "127.0.0.1", "cat"}, bytes.NewReader(want), &out, io.Discard); err != nil { t.Fatalf("transport: %v", err) } if !bytes.Equal(out.Bytes(), want) { t.Fatalf("bridged data mismatch: %q", out.Bytes()) } } func TestTransportRejectsUnknownHostKey(t *testing.T) { keyPath, pub := genClientKey(t) srv := newTestServer(t, pub, func(_ string, _ io.Reader, _, _ io.Writer) int { return 0 }) // Pin a DIFFERENT key, so the server's real host key must be rejected. _, wrong := genClientKey(t) kh := writeKnownHosts(t, "127.0.0.1", srv.port(), wrong) setEnv(t, keyPath, kh, srv.port()) if err := transport([]string{"-l", "u", "127.0.0.1", "true"}, bytes.NewReader(nil), io.Discard, io.Discard); err == nil { t.Fatal("expected host-key mismatch error, got nil") } } func TestTransportPropagatesExitCode(t *testing.T) { keyPath, pub := genClientKey(t) srv := newTestServer(t, pub, func(_ string, _ io.Reader, _, _ io.Writer) int { return 7 }) kh := writeKnownHosts(t, "127.0.0.1", srv.port(), srv.hostKey.PublicKey()) setEnv(t, keyPath, kh, srv.port()) err := transport([]string{"-l", "u", "127.0.0.1", "false"}, bytes.NewReader(nil), io.Discard, io.Discard) var ee *ssh.ExitError if !errors.As(err, &ee) || ee.ExitStatus() != 7 { t.Fatalf("expected exit status 7, got %v", err) } } func TestPubkeyMatchesKeygen(t *testing.T) { keyPath, pub := genClientKey(t) keyData, err := os.ReadFile(keyPath) if err != nil { t.Fatal(err) } t.Setenv("RSH_KEY", "") t.Setenv("RSH_KEY_DATA", string(keyData)) var out bytes.Buffer if err := pubkey(&out); err != nil { t.Fatalf("pubkey: %v", err) } got, _, _, _, err := ssh.ParseAuthorizedKey(out.Bytes()) if err != nil { t.Fatalf("parse pubkey output: %v", err) } if !bytes.Equal(got.Marshal(), pub.Marshal()) { t.Fatal("pubkey output does not match the keygen public key") } } func TestPubkeyRejectsMissingKey(t *testing.T) { t.Setenv("RSH_KEY", "") t.Setenv("RSH_KEY_DATA", "") if err := pubkey(io.Discard); err == nil { t.Fatal("expected error when no key is configured") } } func TestScanPrintsFingerprintAndLine(t *testing.T) { keyPath, pub := genClientKey(t) srv := newTestServer(t, pub, func(_ string, _ io.Reader, _, _ io.Writer) int { return 0 }) t.Setenv("RSH_KEY", keyPath) t.Setenv("RSH_PORT", strconv.Itoa(srv.port())) var out bytes.Buffer if err := scan("u@127.0.0.1", &out); err != nil { t.Fatalf("scan: %v", err) } lines := strings.SplitN(strings.TrimSpace(out.String()), "\n", 2) if len(lines) != 2 { t.Fatalf("want fingerprint + known_hosts line, got %q", out.String()) } if !strings.HasPrefix(lines[0], "SHA256:") { t.Errorf("fingerprint: %q", lines[0]) } // The printed known_hosts line must validate the real host key. khPath := filepath.Join(t.TempDir(), "kh") if err := os.WriteFile(khPath, []byte(lines[1]+"\n"), 0o600); err != nil { t.Fatal(err) } cb, err := knownhosts.New(khPath) if err != nil { t.Fatal(err) } addr := net.JoinHostPort("127.0.0.1", strconv.Itoa(srv.port())) remote := &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: srv.port()} if err := cb(addr, remote, srv.hostKey.PublicKey()); err != nil { t.Errorf("pinned line did not validate the host key: %v", err) } } func TestParseTransport(t *testing.T) { cases := []struct { args []string user, host string cmd string wantErr bool }{ {[]string{"-l", "alice", "host", "rsync", "--server"}, "alice", "host", "rsync --server", false}, {[]string{"bob@host", "echo", "hi"}, "bob", "host", "echo hi", false}, {[]string{"-l", "alice", "carol@host", "x"}, "alice", "host", "x", false}, {[]string{"host"}, "", "host", "", false}, {[]string{"-l"}, "", "", "", true}, {[]string{"-p", "22", "host", "x"}, "", "", "", true}, } for _, c := range cases { u, h, cmd, err := parseTransport(c.args) if c.wantErr { if err == nil { t.Errorf("%v: expected error", c.args) } continue } if err != nil { t.Errorf("%v: %v", c.args, err) continue } if u != c.user || h != c.host || strings.Join(cmd, " ") != c.cmd { t.Errorf("%v: got (%q,%q,%q) want (%q,%q,%q)", c.args, u, h, strings.Join(cmd, " "), c.user, c.host, c.cmd) } } }