From 84e29f70831fc7c03d49d977eecd74a2c39837df Mon Sep 17 00:00:00 2001 From: Chris Roberts <25612130+coderandhiker@users.noreply.github.com> Date: Fri, 18 Sep 2026 18:40:18 +0000 Subject: [PATCH] feat: authenticate sftp installations with ssh keys sftp installations could only authenticate with a password in the url, so servers that disable password authentication could not be managed. sftp now also offers public keys from a running ssh agent (SSH_AUTH_SOCK, or the OpenSSH agent's named pipe on windows) and from the unencrypted default key files in ~/.ssh. Passphrase-protected key files are skipped with a hint to add them to the agent. A password in the url is still tried first, so existing installations behave as before. The test serves sftp in-process with the ssh and sftp packages already in use, so it runs without the docker sftp container. Refs satisfactorymodding/SatisfactoryModManager#304 --- cli/disk/sftp.go | 124 ++++++++++++++++++++++++++-- cli/disk/sftp_test.go | 183 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 302 insertions(+), 5 deletions(-) create mode 100644 cli/disk/sftp_test.go diff --git a/cli/disk/sftp.go b/cli/disk/sftp.go index 8cfaba8..17bea36 100644 --- a/cli/disk/sftp.go +++ b/cli/disk/sftp.go @@ -5,12 +5,17 @@ import ( "errors" "fmt" "io" + "io/fs" "log/slog" + "net" "net/url" "os" + "path/filepath" + "runtime" "github.com/pkg/sftp" "golang.org/x/crypto/ssh" + "golang.org/x/crypto/ssh/agent" ) var _ Disk = (*sftpDisk)(nil) @@ -38,15 +43,17 @@ func newSFTP(path string) (Disk, error) { return nil, fmt.Errorf("failed to parse sftp url: %w", err) } - password, ok := u.User.Password() - var auth []ssh.AuthMethod - if ok { - auth = append(auth, ssh.Password(password)) + agentConn, err := dialAgent() + if err != nil { + slog.Warn("not using ssh agent", slog.Any("err", err)) + } + if agentConn != nil { + defer agentConn.Close() } conn, err := ssh.Dial("tcp", u.Host, &ssh.ClientConfig{ User: u.User.Username(), - Auth: auth, + Auth: sshAuthMethods(u, agentConn), // TODO Somehow use systems hosts file HostKeyCallback: ssh.InsecureIgnoreHostKey(), @@ -167,3 +174,110 @@ func (l sftpDisk) Open(path string, _ int) (io.WriteCloser, error) { return f, nil } + +// defaultKeyFiles are the private keys tried from ~/.ssh, in the order ssh tries them +var defaultKeyFiles = []string{"id_ed25519", "id_ecdsa", "id_rsa"} + +// sshAuthMethods returns the auth methods for u: the password in the url if there is one, +// then public keys from the agent (if agentConn is not nil) and from the default key files +func sshAuthMethods(u *url.URL, agentConn io.ReadWriter) []ssh.AuthMethod { + var auth []ssh.AuthMethod + if password, ok := u.User.Password(); ok { + auth = append(auth, ssh.Password(password)) + } + + var agentClient agent.ExtendedAgent + if agentConn != nil { + agentClient = agent.NewClient(agentConn) + } + fileSigners := loadDefaultKeys() + + // ssh tries each auth method name only once, so agent and file keys have to share one publickey method + return append(auth, ssh.PublicKeysCallback(func() ([]ssh.Signer, error) { + var signers []ssh.Signer + if agentClient != nil { + agentSigners, err := agentClient.Signers() + if err != nil { + slog.Warn("failed to list ssh agent keys", slog.Any("err", err)) + } + signers = append(signers, agentSigners...) + } + return dedupeSigners(append(signers, fileSigners...)), nil + })) +} + +// dialAgent connects to the ssh agent at SSH_AUTH_SOCK, or on windows to the OpenSSH agent's named pipe. +// It returns nil, nil when there is no agent +func dialAgent() (io.ReadWriteCloser, error) { + sock := os.Getenv("SSH_AUTH_SOCK") + + if runtime.GOOS == "windows" { + if sock == "" { + sock = `\\.\pipe\openssh-ssh-agent` + } + // The pipe opens like a file, and the agent protocol is request/response, so plain file I/O is enough + f, err := os.OpenFile(sock, os.O_RDWR, 0) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + return nil, nil + } + return nil, fmt.Errorf("failed to open ssh agent pipe %s: %w", sock, err) + } + return f, nil + } + + if sock == "" { + return nil, nil + } + conn, err := net.Dial("unix", sock) + if err != nil { + return nil, fmt.Errorf("failed to connect to ssh agent at %s: %w", sock, err) + } + return conn, nil +} + +// loadDefaultKeys returns signers for the unencrypted default key files in ~/.ssh +func loadDefaultKeys() []ssh.Signer { + home, err := os.UserHomeDir() + if err != nil { + return nil + } + + signers := make([]ssh.Signer, 0, len(defaultKeyFiles)) + for _, name := range defaultKeyFiles { + keyPath := filepath.Join(home, ".ssh", name) + pem, err := os.ReadFile(keyPath) + if err != nil { + continue + } + + signer, err := ssh.ParsePrivateKey(pem) + if err != nil { + var missing *ssh.PassphraseMissingError + if errors.As(err, &missing) { + slog.Info("skipping passphrase-protected ssh key, add it to the ssh agent to use it", slog.String("path", keyPath)) + } else { + slog.Warn("failed to parse ssh key", slog.String("path", keyPath), slog.Any("err", err)) + } + continue + } + signers = append(signers, signer) + } + return signers +} + +// dedupeSigners drops repeated keys, so a key that is both in the agent and on disk +// only counts once against the server's MaxAuthTries +func dedupeSigners(signers []ssh.Signer) []ssh.Signer { + seen := make(map[string]bool, len(signers)) + unique := signers[:0] + for _, s := range signers { + key := string(s.PublicKey().Marshal()) + if seen[key] { + continue + } + seen[key] = true + unique = append(unique, s) + } + return unique +} diff --git a/cli/disk/sftp_test.go b/cli/disk/sftp_test.go new file mode 100644 index 0000000..98437c2 --- /dev/null +++ b/cli/disk/sftp_test.go @@ -0,0 +1,183 @@ +package disk + +import ( + "bytes" + "crypto/ed25519" + "crypto/rand" + "crypto/x509" + "encoding/pem" + "errors" + "net" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/MarvinJWendt/testza" + "github.com/pkg/sftp" + "golang.org/x/crypto/ssh" + "golang.org/x/crypto/ssh/agent" +) + +func TestSFTPAuth(t *testing.T) { + key, signer := newTestKey(t) + addr := startSFTPServer(t, signer.PublicKey()) + dir := t.TempDir() + + // An empty home and no agent, so only what each case sets up can authenticate + home := t.TempDir() + t.Setenv("HOME", home) + t.Setenv("USERPROFILE", home) + t.Setenv("SSH_AUTH_SOCK", "") + + t.Run("password", func(t *testing.T) { + d, err := newSFTP("sftp://user:pass@" + addr + "/") + testza.AssertNoError(t, err) + assertSFTPWorks(t, d, dir) + }) + + t.Run("no credentials", func(t *testing.T) { + _, err := newSFTP("sftp://user@" + addr + "/") + testza.AssertNotNil(t, err) + }) + + t.Run("key file", func(t *testing.T) { + writeKeyFile(t, filepath.Join(home, ".ssh", "id_ed25519"), key) + defer os.RemoveAll(filepath.Join(home, ".ssh")) + + d, err := newSFTP("sftp://user@" + addr + "/") + testza.AssertNoError(t, err) + assertSFTPWorks(t, d, dir) + }) + + t.Run("agent", func(t *testing.T) { + if runtime.GOOS == "windows" { + // The agent is a named pipe on windows, which this test cannot serve + return + } + + keyring := agent.NewKeyring() + testza.AssertNoError(t, keyring.Add(agent.AddedKey{PrivateKey: key})) + t.Setenv("SSH_AUTH_SOCK", startAgent(t, keyring)) + + d, err := newSFTP("sftp://user@" + addr + "/") + testza.AssertNoError(t, err) + assertSFTPWorks(t, d, dir) + }) +} + +func TestDedupeSigners(t *testing.T) { + _, a := newTestKey(t) + _, b := newTestKey(t) + testza.AssertEqual(t, []ssh.Signer{a, b}, dedupeSigners([]ssh.Signer{a, b, a, b, a})) +} + +func assertSFTPWorks(t *testing.T, d Disk, dir string) { + file := filepath.Join(dir, "hello.txt") + testza.AssertNoError(t, d.Write(file, []byte("hello"))) + + data, err := d.Read(file) + testza.AssertNoError(t, err) + testza.AssertEqual(t, "hello", string(data)) +} + +func newTestKey(t *testing.T) (ed25519.PrivateKey, ssh.Signer) { + _, key, err := ed25519.GenerateKey(rand.Reader) + testza.AssertNoError(t, err) + signer, err := ssh.NewSignerFromKey(key) + testza.AssertNoError(t, err) + return key, signer +} + +func writeKeyFile(t *testing.T, path string, key ed25519.PrivateKey) { + der, err := x509.MarshalPKCS8PrivateKey(key) + testza.AssertNoError(t, err) + testza.AssertNoError(t, os.MkdirAll(filepath.Dir(path), 0o700)) + testza.AssertNoError(t, os.WriteFile(path, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: der}), 0o600)) +} + +// startAgent serves keyring on a unix socket and returns the socket path +func startAgent(t *testing.T, keyring agent.Agent) string { + sock := filepath.Join(t.TempDir(), "agent.sock") + listener, err := net.Listen("unix", sock) + testza.AssertNoError(t, err) + t.Cleanup(func() { listener.Close() }) + + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + go func() { _ = agent.ServeAgent(keyring, conn) }() + } + }() + + return sock +} + +// startSFTPServer serves sftp on a random localhost port, accepting user:pass and the authorized key, and returns its address +func startSFTPServer(t *testing.T, authorized ssh.PublicKey) string { + _, hostKey := newTestKey(t) + config := &ssh.ServerConfig{ + PasswordCallback: func(c ssh.ConnMetadata, password []byte) (*ssh.Permissions, error) { + if c.User() == "user" && string(password) == "pass" { + return nil, nil + } + return nil, errors.New("wrong password") + }, + PublicKeyCallback: func(c ssh.ConnMetadata, key ssh.PublicKey) (*ssh.Permissions, error) { + if c.User() == "user" && bytes.Equal(key.Marshal(), authorized.Marshal()) { + return nil, nil + } + return nil, errors.New("unknown key") + }, + } + config.AddHostKey(hostKey) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + testza.AssertNoError(t, err) + t.Cleanup(func() { listener.Close() }) + + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + go serveSFTP(conn, config) + } + }() + + return listener.Addr().String() +} + +func serveSFTP(conn net.Conn, config *ssh.ServerConfig) { + sshConn, channels, requests, err := ssh.NewServerConn(conn, config) + if err != nil { + return + } + defer sshConn.Close() + go ssh.DiscardRequests(requests) + + for newChannel := range channels { + channel, channelRequests, err := newChannel.Accept() + if err != nil { + return + } + go func() { + defer channel.Close() + for req := range channelRequests { + // The payload is an ssh string: a length prefix followed by the subsystem name + isSFTP := req.Type == "subsystem" && len(req.Payload) > 4 && string(req.Payload[4:]) == "sftp" + _ = req.Reply(isSFTP, nil) + if isSFTP { + if server, err := sftp.NewServer(channel); err == nil { + _ = server.Serve() + } + return + } + } + }() + } +}