diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2025-04-10 07:50:41 -0700 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2025-04-10 10:51:03 -0700 |
| commit | 89b48b3c0d4ed4eceebd76d4fd3d981b642d3824 (patch) | |
| tree | bec39fb4e71c87ccd6c9b2b9d8aff3670056f3cd /internal/wfs | |
| parent | Copy recursive (diff) | |
| download | ccp-89b48b3c0d4ed4eceebd76d4fd3d981b642d3824.tar.zst | |
Add support for SFTP remote file copies
Diffstat (limited to 'internal/wfs')
| -rw-r--r-- | internal/wfs/osfs/osfs.go | 58 | ||||
| -rw-r--r-- | internal/wfs/sftpfs/sftpfs.go | 234 | ||||
| -rw-r--r-- | internal/wfs/wfs.go | 86 |
3 files changed, 378 insertions, 0 deletions
diff --git a/internal/wfs/osfs/osfs.go b/internal/wfs/osfs/osfs.go new file mode 100644 index 0000000..bf6951c --- /dev/null +++ b/internal/wfs/osfs/osfs.go @@ -0,0 +1,58 @@ +package osfs + +import ( + "io" + "io/fs" + "os" + + "gitlab.com/rhogenson/ccp/internal/wfs" +) + +var ( + _ wfs.FS = FS{} + _ wfs.MkdirModeFS = FS{} + _ wfs.ReadLinkFS = FS{} + _ fs.StatFS = FS{} +) + +type FS struct{} + +func (FS) Open(name string) (fs.File, error) { + return os.Open(name) +} + +func (FS) Stat(name string) (fs.FileInfo, error) { + return os.Stat(name) +} + +func (FS) Lstat(name string) (fs.FileInfo, error) { + return os.Lstat(name) +} + +func (FS) ReadLink(name string) (string, error) { + return os.Readlink(name) +} + +func (FS) Create(name string, perm fs.FileMode) (io.WriteCloser, error) { + return os.OpenFile(name, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, perm) +} + +func (FS) Remove(name string) error { + return os.Remove(name) +} + +func (FS) Mkdir(name string) error { + return os.Mkdir(name, 0700) +} + +func (FS) MkdirMode(name string, mode fs.FileMode) error { + return os.Mkdir(name, mode) +} + +func (FS) Symlink(oldname, newname string) error { + return os.Symlink(oldname, newname) +} + +func (FS) Chmod(name string, mode fs.FileMode) error { + return os.Chmod(name, mode) +} diff --git a/internal/wfs/sftpfs/sftpfs.go b/internal/wfs/sftpfs/sftpfs.go new file mode 100644 index 0000000..1457475 --- /dev/null +++ b/internal/wfs/sftpfs/sftpfs.go @@ -0,0 +1,234 @@ +package sftpfs + +import ( + "errors" + "fmt" + "io" + "io/fs" + "net" + "os" + "path" + "path/filepath" + "strings" + "sync" + + "github.com/pkg/sftp" + "gitlab.com/rhogenson/ccp/internal/wfs" + "golang.org/x/crypto/ssh" + "golang.org/x/crypto/ssh/agent" + "golang.org/x/crypto/ssh/knownhosts" + "golang.org/x/term" +) + +var ( + _ wfs.FS = (*FS)(nil) + _ wfs.ReadLinkFS = (*FS)(nil) + _ fs.StatFS = (*FS)(nil) + _ fs.StatFS = (*FS)(nil) + _ fs.ReadDirFS = (*FS)(nil) +) + +type FS struct { + User, Host string + conn *sftp.Client + sshConn *ssh.Client +} + +var sshAgent = sync.OnceValue(func() agent.ExtendedAgent { + socket := os.Getenv("SSH_AUTH_SOCK") + if socket == "" { + return nil + } + conn, err := net.Dial("unix", socket) + if err != nil { + return nil + } + return agent.NewClient(conn) +}) + +func sshKeys() ([]ssh.Signer, error) { + sshAgent := sshAgent() + if sshAgent != nil { + if signers, err := sshAgent.Signers(); err == nil && len(signers) > 0 { + return signers, nil + } + } + sshDir := filepath.Join(os.Getenv("HOME"), ".ssh") + sshFiles, err := os.ReadDir(sshDir) + if len(sshFiles) == 0 { + return nil, err + } + var keys []ssh.Signer + var passwordProtectedKey []byte + var passwordProtectedKeyFile string + for _, f := range sshFiles { + if f.Name() == "known_hosts" || strings.HasSuffix(f.Name(), ".pub") { + continue + } + fileName := filepath.Join(sshDir, f.Name()) + keyBytes, err := os.ReadFile(fileName) + if err != nil { + continue + } + key, err := ssh.ParsePrivateKey(keyBytes) + if err != nil { + if passwordProtectedKey == nil && errors.As(err, new(*ssh.PassphraseMissingError)) { + passwordProtectedKey = keyBytes + passwordProtectedKeyFile = fileName + } + continue + } + keys = append(keys, key) + } + if len(keys) == 0 && passwordProtectedKey != nil { + fmt.Fprintf(os.Stderr, "Enter password for %s: ", passwordProtectedKeyFile) + for i := range 3 { + if i > 0 { + fmt.Fprintf(os.Stderr, "Incorrect password, try again: ") + } + password, err := term.ReadPassword(int(os.Stdin.Fd())) + fmt.Fprintln(os.Stderr) + if err != nil { + return nil, err + } + key, err := ssh.ParseRawPrivateKeyWithPassphrase(passwordProtectedKey, password) + if err != nil { + continue + } + if sshAgent != nil { + sshAgent.Add(agent.AddedKey{PrivateKey: key}) + } + signer, err := ssh.NewSignerFromKey(key) + if err != nil { + return nil, err + } + return []ssh.Signer{signer}, nil + } + return nil, errors.New("user couldn't remember her password") + } + return keys, nil +} + +func appendToKnownHosts(hostname string, key ssh.PublicKey) error { + f, err := os.OpenFile(path.Join(os.Getenv("HOME"), ".ssh/known_hosts"), os.O_WRONLY|os.O_CREATE|os.O_APPEND, 0600) + if err != nil { + return err + } + defer f.Close() + if _, err := f.WriteString(knownhosts.Line([]string{hostname}, key) + "\n"); err != nil { + return err + } + return f.Close() +} + +func Dial(target string) (*FS, error) { + knownHostChecker, err := knownhosts.New(path.Join(os.Getenv("HOME"), ".ssh/known_hosts")) + if err != nil { + knownHostChecker = func(string, net.Addr, ssh.PublicKey) error { return &knownhosts.KeyError{} } + } + var user string + if i := strings.Index(target, "@"); i >= 0 { + user, target = target[:i], target[i+1:] + } else { + user = os.Getenv("USER") + } + sshConn, err := ssh.Dial("tcp", target+":22", &ssh.ClientConfig{ + User: user, + Auth: []ssh.AuthMethod{ + ssh.PublicKeysCallback(sshKeys), + ssh.RetryableAuthMethod(ssh.PasswordCallback(func() (string, error) { + fmt.Fprintf(os.Stderr, "Enter password for %s@%s: ", user, target) + password, err := term.ReadPassword(int(os.Stdin.Fd())) + fmt.Fprintln(os.Stderr) + return string(password), err + }), 3), + }, + HostKeyCallback: func(hostname string, remote net.Addr, key ssh.PublicKey) error { + err := knownHostChecker(hostname, remote, key) + if err == nil { + return nil + } + var keyErr *knownhosts.KeyError + if !errors.As(err, &keyErr) || len(keyErr.Want) > 0 { + return err + } + appendToKnownHosts(hostname, key) + return nil + }, + }) + if err != nil { + return nil, err + } + sftpConn, err := sftp.NewClient(sshConn) + if err != nil { + sshConn.Close() + return nil, err + } + return &FS{ + User: user, + Host: target, + conn: sftpConn, + sshConn: sshConn, + }, nil +} + +func (f *FS) Close() error { + sftpErr := f.conn.Close() + if err := f.sshConn.Close(); err != nil { + return err + } + return sftpErr +} + +func (f *FS) Open(name string) (fs.File, error) { + return f.conn.Open(name) +} + +func (f *FS) ReadDir(name string) ([]fs.DirEntry, error) { + entriesFileInfo, err := f.conn.ReadDir(name) + entries := make([]fs.DirEntry, len(entriesFileInfo)) + for i, entry := range entriesFileInfo { + entries[i] = fs.FileInfoToDirEntry(entry) + } + return entries, err +} + +func (f *FS) Stat(name string) (fs.FileInfo, error) { + return f.conn.Stat(name) +} + +func (f *FS) Lstat(name string) (fs.FileInfo, error) { + return f.conn.Lstat(name) +} + +func (f *FS) ReadLink(name string) (string, error) { + return f.conn.ReadLink(name) +} + +func (f *FS) Create(name string, perm fs.FileMode) (io.WriteCloser, error) { + file, err := f.conn.Create(name) + if err != nil { + return nil, err + } + if err := file.Chmod(perm); err != nil { + file.Close() + return nil, err + } + return file, nil +} + +func (f *FS) Remove(name string) error { + return f.conn.Remove(name) +} + +func (f *FS) Mkdir(name string) error { + return f.conn.Mkdir(name) +} + +func (f *FS) Symlink(oldname, newname string) error { + return f.conn.Symlink(oldname, newname) +} + +func (f *FS) Chmod(name string, mode fs.FileMode) error { + return f.conn.Chmod(name, mode) +} diff --git a/internal/wfs/wfs.go b/internal/wfs/wfs.go new file mode 100644 index 0000000..bcdca00 --- /dev/null +++ b/internal/wfs/wfs.go @@ -0,0 +1,86 @@ +package wfs + +import ( + "io" + "io/fs" + "path" +) + +type ReadLinkFS interface { + fs.FS + + ReadLink(string) (string, error) + Lstat(string) (fs.FileInfo, error) +} + +func ReadLink(fsys fs.FS, name string) (string, error) { + sym, ok := fsys.(ReadLinkFS) + if !ok { + return "", &fs.PathError{Op: "readlink", Path: name, Err: fs.ErrInvalid} + } + return sym.ReadLink(name) +} + +func Lstat(fsys fs.FS, name string) (fs.FileInfo, error) { + sym, ok := fsys.(ReadLinkFS) + if !ok { + return fs.Stat(fsys, name) + } + return sym.Lstat(name) +} + +type FS interface { + fs.FS + + Create(string, fs.FileMode) (io.WriteCloser, error) + Remove(string) error + Mkdir(string) error + Symlink(string, string) error + Chmod(string, fs.FileMode) error +} + +type MkdirModeFS interface { + FS + + MkdirMode(string, fs.FileMode) error +} + +func MkdirMode(fsys FS, name string, mode fs.FileMode) error { + if fsys, ok := fsys.(MkdirModeFS); ok { + return fsys.MkdirMode(name, mode) + } + if err := fsys.Mkdir(name); err != nil { + return err + } + return fsys.Chmod(name, mode) +} + +func removeDir(fsys FS, dir string) error { + entries, err := fs.ReadDir(fsys, dir) + if err != nil { + return err + } + for _, f := range entries { + name := path.Join(dir, f.Name()) + if f.IsDir() { + err = removeDir(fsys, name) + } else { + err = fsys.Remove(name) + } + if err != nil { + return err + } + } + return fsys.Remove(dir) +} + +func RemoveAll(fsys FS, path string) error { + stat, err := Lstat(fsys, path) + if err != nil { + return err + } + if stat.IsDir() { + return removeDir(fsys, path) + } + return fsys.Remove(path) +} |
