aboutsummaryrefslogtreecommitdiffstats
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/cp/cp.go237
-rw-r--r--internal/wfs/osfs/osfs.go58
-rw-r--r--internal/wfs/sftpfs/sftpfs.go234
-rw-r--r--internal/wfs/wfs.go86
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)
+}