diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/cp/cp.go | 237 | ||||
| -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 |
4 files changed, 529 insertions, 86 deletions
diff --git a/internal/cp/cp.go b/internal/cp/cp.go index c19627f..d944f9c 100644 --- a/internal/cp/cp.go +++ b/internal/cp/cp.go @@ -5,21 +5,36 @@ import ( "fmt" "io" "io/fs" - "os" - "path/filepath" + "path" + "sync" + + "gitlab.com/rhogenson/ccp/internal/wfs" + "gitlab.com/rhogenson/ccp/internal/wfs/sftpfs" ) type Progress interface { Max(int64) Progress(int64) - FileStart(string) + FileStart(string, string) FileDone(string, error) } -func size(files []string) int64 { +type FSPath struct { + FS wfs.FS + Path string +} + +func (p FSPath) String() string { + if fsys, ok := p.FS.(*sftpfs.FS); ok { + return fsys.User + "@" + fsys.Host + ":" + p.Path + } + return p.Path +} + +func size(srcs []FSPath) int64 { var n int64 = 0 - for _, f := range files { - filepath.WalkDir(f, func(path string, d fs.DirEntry, err error) error { + for _, src := range srcs { + fs.WalkDir(src.FS, src.Path, func(_ string, d fs.DirEntry, err error) error { if err != nil { return nil } @@ -30,7 +45,7 @@ func size(files []string) int64 { return nil } n += 1 + stat.Size() - case os.ModeSymlink, os.ModeDir: + case fs.ModeSymlink, fs.ModeDir: n++ } return nil @@ -39,15 +54,40 @@ func size(files []string) int64 { return n } -func fileExists(path string) bool { - _, err := os.Lstat(path) - return !errors.Is(err, os.ErrNotExist) +func fileExists(path FSPath) bool { + _, err := wfs.Lstat(path.FS, path.Path) + return !errors.Is(err, fs.ErrNotExist) +} + +type copier struct { + p Progress + sem chan struct{} + + force bool +} + +func (c *copier) g(fn func()) { + c.sem <- struct{}{} + go func() { + defer func() { <-c.sem }() + fn() + }() +} + +func (c *copier) openWithRetry(path FSPath, fn func() error) error { + if err := fn(); err == nil || !c.force || !fileExists(path) { + return err + } + if err := wfs.RemoveAll(path.FS, path.Path); err != nil { + return err + } + return fn() } -func copyRegularFile(src, dst string, force bool, progress Progress) (err error) { - progress.FileStart(src) +func (c *copier) copyRegularFile(src, dst FSPath) error { + c.p.FileStart(src.String(), dst.String()) - in, err := os.Open(src) + in, err := src.FS.Open(src.Path) if err != nil { return err } @@ -56,117 +96,142 @@ func copyRegularFile(src, dst string, force bool, progress Progress) (err error) if err != nil { return err } - out, err := os.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, stat.Mode().Perm()) - if err != nil { - if !force || !fileExists(dst) { - return err - } - if err := os.RemoveAll(dst); err != nil { - return err - } - out, err = os.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, stat.Mode().Perm()) - if err != nil { - return err - } + var out io.WriteCloser + if err := c.openWithRetry(dst, func() error { + var err error + out, err = dst.FS.Create(dst.Path, stat.Mode().Perm()) + return err + }); err != nil { + return err } - defer out.Close() for { n, err := io.CopyN(out, in, 1024*1024) if n > 0 { - progress.Progress(n) + c.p.Progress(n) } if err != nil { if err == io.EOF { break } + out.Close() return err } } if err := out.Close(); err != nil { return err } - progress.Progress(1) + c.p.Progress(1) + c.p.FileDone(src.String(), nil) return nil } -func copySpecialFile(src string, d fs.DirEntry, dst string, force bool, progress Progress) error { - switch d.Type() { - case fs.ModeSymlink: - target, err := os.Readlink(src) - if err != nil { - return err +func (c *copier) copySymlink(src FSPath, d fs.DirEntry, dst FSPath) error { + target, err := wfs.ReadLink(src.FS, src.Path) + if err != nil { + return err + } + if err := c.openWithRetry(dst, func() error { + return dst.FS.Symlink(target, dst.Path) + }); err != nil { + return err + } + c.p.Progress(1) + return nil +} + +func (c *copier) copyDir(src FSPath, d fs.DirEntry, dst FSPath) error { + stat, err := d.Info() + if err != nil { + return err + } + hasWritePerm := stat.Mode()&0300 == 0300 + if err := c.openWithRetry(dst, func() error { + if hasWritePerm { + return wfs.MkdirMode(dst.FS, dst.Path, stat.Mode().Perm()) + } else { + return dst.FS.Mkdir(dst.Path) } - if err := os.Symlink(target, dst); err != nil { - if !force || !fileExists(dst) { - return err - } - if err := os.RemoveAll(dst); err != nil { - return err - } - if err := os.Symlink(target, dst); err != nil { - return err + }); err != nil { + return err + } + entries, err := fs.ReadDir(src.FS, src.Path) + wg := new(sync.WaitGroup) + <-c.sem + for _, entry := range entries { + wg.Add(1) + c.g(func() { + defer wg.Done() + src := FSPath{src.FS, path.Join(src.Path, entry.Name())} + if err := c.copyFile(src, entry, FSPath{dst.FS, path.Join(dst.Path, entry.Name())}); err != nil { + c.p.FileDone(src.String(), err) } - } - progress.Progress(1) - case fs.ModeDir: - stat, err := d.Info() - if err != nil { + }) + } + wg.Wait() + c.sem <- struct{}{} + if err != nil { + return err + } + if !hasWritePerm { + if err := dst.FS.Chmod(dst.Path, stat.Mode().Perm()); err != nil { return err } - if err := os.Mkdir(dst, stat.Mode().Perm()); err != nil { - if !force || !fileExists(dst) { - return err - } - if err := os.RemoveAll(dst); err != nil { - return err - } - if err := os.Mkdir(dst, stat.Mode().Perm()); err != nil { - return err - } - } - progress.Progress(1) + } + c.p.Progress(1) + return nil +} + +func (c *copier) copyFile(src FSPath, d fs.DirEntry, dst FSPath) error { + switch d.Type() { + case 0: // regular file + return c.copyRegularFile(src, dst) + case fs.ModeDir: + return c.copyDir(src, d, dst) + case fs.ModeSymlink: + return c.copySymlink(src, d, dst) default: return fmt.Errorf("%s: unknown file type %s", src, d.Type()) } - return nil } -func Copy(srcs []string, dstRoot string, force bool, progress Progress) { +func Copy(progress Progress, srcs []FSPath, dstRoot FSPath, force bool) { go func() { progress.Max(size(srcs)) }() + + dstIsDir := true + if len(srcs) == 1 { + stat, err := fs.Stat(dstRoot.FS, dstRoot.Path) + dstIsDir = err == nil && stat.IsDir() + } + const maxConcurrency = 10 // sem acts as a semaphore to limit the number of concurrent file copies sem := make(chan struct{}, maxConcurrency) + c := &copier{ + p: progress, + sem: sem, + force: force, + } + wg := new(sync.WaitGroup) for _, srcRoot := range srcs { - filepath.WalkDir(srcRoot, func(src string, d fs.DirEntry, err error) error { - if err != nil { - progress.FileDone(src, err) - return nil + wg.Add(1) + c.g(func() { + defer wg.Done() + dstRoot := dstRoot + if dstIsDir { + dstRoot.Path = path.Join(dstRoot.Path, path.Base(srcRoot.Path)) } - relPath, err := filepath.Rel(srcRoot, src) + stat, err := fs.Stat(srcRoot.FS, srcRoot.Path) if err != nil { - progress.FileDone(src, err) - return nil + progress.FileDone(srcRoot.String(), err) + return } - dst := filepath.Join(dstRoot, relPath) - if d.Type().IsRegular() { - sem <- struct{}{} - go func() { - defer func() { <-sem }() - err := copyRegularFile(src, dst, force, progress) - progress.FileDone(src, err) - }() - } else { - if err := copySpecialFile(src, d, dst, force, progress); err != nil { - progress.FileDone(src, err) - return nil - } + if err := c.copyFile(srcRoot, fs.FileInfoToDirEntry(stat), dstRoot); err != nil { + progress.FileDone(srcRoot.String(), err) + return } - return nil }) } - for range maxConcurrency { - sem <- struct{}{} - } + wg.Wait() } 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) +} |
