Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 119 additions & 5 deletions cli/disk/sftp.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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
}
183 changes: 183 additions & 0 deletions cli/disk/sftp_test.go
Original file line number Diff line number Diff line change
@@ -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
}
}
}()
}
}