From 72630d1df522ce12debe157ada9196359d831e29 Mon Sep 17 00:00:00 2001 From: a-palchikov Date: Wed, 6 Jan 2021 13:21:06 +0100 Subject: [PATCH] Implement support for preserving file times for 'tsh scp' (#4764) * Add -p flag to scp * Add support for preserving access/modification times on files/directories when copying files between hosts. * lib/sshutils/scp: add time statting for directories * Add directory handling for scp * Rewrite scp tests with testify * Address review comments --- lib/client/api.go | 24 +- lib/client/client.go | 9 +- lib/sshutils/scp/http.go | 27 +- lib/sshutils/scp/local.go | 31 +- lib/sshutils/scp/scp.go | 286 +++++++---- lib/sshutils/scp/scp_test.go | 782 +++++++++++++++++++++---------- lib/sshutils/scp/stat_darwin.go | 32 ++ lib/sshutils/scp/stat_linux.go | 32 ++ lib/sshutils/scp/stat_windows.go | 28 ++ lib/web/files.go | 5 +- tool/teleport/common/teleport.go | 1 + tool/tsh/tsh.go | 11 +- 12 files changed, 911 insertions(+), 357 deletions(-) create mode 100644 lib/sshutils/scp/stat_darwin.go create mode 100644 lib/sshutils/scp/stat_linux.go create mode 100644 lib/sshutils/scp/stat_windows.go diff --git a/lib/client/api.go b/lib/client/api.go index 3c6f2761c1a..37a427a6bf4 100644 --- a/lib/client/api.go +++ b/lib/client/api.go @@ -1329,9 +1329,9 @@ func (tc *TeleportClient) ExecuteSCP(ctx context.Context, cmd scp.Command) (err } // SCP securely copies file(s) from one SSH server to another -func (tc *TeleportClient) SCP(ctx context.Context, args []string, port int, recursive bool, quiet bool) (err error) { +func (tc *TeleportClient) SCP(ctx context.Context, args []string, port int, flags scp.Flags, quiet bool) (err error) { if len(args) < 2 { - return trace.Errorf("Need at least two arguments for scp") + return trace.Errorf("need at least two arguments for scp") } first := args[0] last := args[len(args)-1] @@ -1344,7 +1344,7 @@ func (tc *TeleportClient) SCP(ctx context.Context, args []string, port int, recu if !tc.Config.ProxySpecified() { return trace.BadParameter("proxy server is not specified") } - log.Infof("Connecting to proxy to copy (recursively=%v)...", recursive) + log.Infof("Connecting to proxy to copy (recursively=%v)...", flags.Recursive) proxyClient, err := tc.ConnectToProxy(ctx) if err != nil { return trace.Wrap(err) @@ -1407,12 +1407,10 @@ func (tc *TeleportClient) SCP(ctx context.Context, args []string, port int, recu User: tc.Username, ProgressWriter: progressWriter, RemoteLocation: dest.Path, - Flags: scp.Flags{ - Target: []string{src}, - Recursive: recursive, - DirectoryMode: directoryMode, - }, + Flags: flags, } + scpConfig.Flags.Target = []string{src} + scpConfig.Flags.DirectoryMode = directoryMode cmd, err := scp.CreateUploadCommand(scpConfig) if err != nil { @@ -1424,8 +1422,8 @@ func (tc *TeleportClient) SCP(ctx context.Context, args []string, port int, recu return onError(err) } } - // download: } else { + // download: src, err := scp.ParseSCPDestination(first) if err != nil { return trace.Wrap(err) @@ -1441,14 +1439,12 @@ func (tc *TeleportClient) SCP(ctx context.Context, args []string, port int, recu // copy everything except the last arg (that's destination) for _, dest := range args[1:] { scpConfig := scp.Config{ - User: tc.Username, - Flags: scp.Flags{ - Recursive: recursive, - Target: []string{dest}, - }, + User: tc.Username, + Flags: flags, RemoteLocation: src.Path, ProgressWriter: progressWriter, } + scpConfig.Flags.Target = []string{dest} cmd, err := scp.CreateDownloadCommand(scpConfig) if err != nil { diff --git a/lib/client/client.go b/lib/client/client.go index 0a596b51778..91d9fb41235 100644 --- a/lib/client/client.go +++ b/lib/client/client.go @@ -876,17 +876,18 @@ func (c *NodeClient) ExecuteSCP(cmd scp.Command) error { &net.IPAddr{}, ) - closeC := make(chan interface{}, 1) + closeC := make(chan error, 1) go func() { - if err = cmd.Execute(ch); err != nil { + err := cmd.Execute(ch) + if err != nil { log.Error(err) } stdin.Close() - close(closeC) + closeC <- err }() runErr := s.Run(shellCmd) - <-closeC + err = <-closeC if runErr != nil && (err == nil || trace.IsEOF(err)) { err = runErr diff --git a/lib/sshutils/scp/http.go b/lib/sshutils/scp/http.go index 1e5adc91dbf..8e0220810f2 100644 --- a/lib/sshutils/scp/http.go +++ b/lib/sshutils/scp/http.go @@ -24,6 +24,7 @@ import ( "os" "path/filepath" "strconv" + "time" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/httplib" @@ -65,7 +66,7 @@ func (r *HTTPTransferRequest) parseRemoteLocation() (string, string, error) { return dir, filename, nil } -// CreateHTTPUpload creates HTTP download command +// CreateHTTPUpload creates an HTTP upload command func CreateHTTPUpload(req HTTPTransferRequest) (Command, error) { if req.HTTPRequest == nil { return nil, trace.BadParameter("missing parameter HTTPRequest") @@ -113,7 +114,7 @@ func CreateHTTPUpload(req HTTPTransferRequest) (Command, error) { return cmd, nil } -// CreateHTTPDownload creates HTTP upload command +// CreateHTTPDownload creates an HTTP download command func CreateHTTPDownload(req HTTPTransferRequest) (Command, error) { _, filename, err := req.parseRemoteLocation() if err != nil { @@ -150,9 +151,15 @@ type httpFileSystem struct { fileSize int64 } -// SetChmod sets file permissions. It does nothing as there are no permissions +// Chmod sets file permissions. It does nothing as there are no permissions // while processing HTTP downloads -func (l *httpFileSystem) SetChmod(path string, mode int) error { +func (l *httpFileSystem) Chmod(path string, mode int) error { + return nil +} + +// Chtimes sets file access and modification time. +// It is a no-op for the HTTP file system implementation +func (l *httpFileSystem) Chtimes(path string, atime, mtime time.Time) error { return nil } @@ -243,6 +250,18 @@ func (l *httpFileInfo) GetModePerm() os.FileMode { return httpUploadFileMode } +// GetModTime returns file modification time. +// It is a no-op for HTTP file information +func (l *httpFileInfo) GetModTime() time.Time { + return time.Time{} +} + +// GetAccessTime returns file last access time. +// It is a no-op for HTTP file information +func (l *httpFileInfo) GetAccessTime() time.Time { + return time.Time{} +} + type nopWriteCloser struct { io.Writer } diff --git a/lib/sshutils/scp/local.go b/lib/sshutils/scp/local.go index 484e9d3dab8..f53cf9e7bf2 100644 --- a/lib/sshutils/scp/local.go +++ b/lib/sshutils/scp/local.go @@ -20,6 +20,7 @@ import ( "io" "os" "path/filepath" + "time" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" @@ -30,8 +31,8 @@ import ( type localFileSystem struct { } -// SetChmod sets file permissions -func (l *localFileSystem) SetChmod(path string, mode int) error { +// Chmod sets file permissions +func (l *localFileSystem) Chmod(path string, mode int) error { chmode := os.FileMode(mode & int(os.ModePerm)) if err := os.Chmod(path, chmode); err != nil { return trace.Wrap(err) @@ -40,6 +41,11 @@ func (l *localFileSystem) SetChmod(path string, mode int) error { return nil } +// Chtimes sets file access and modification times +func (l *localFileSystem) Chtimes(path string, atime, mtime time.Time) error { + return trace.ConvertSystemError(os.Chtimes(path, atime, mtime)) +} + // MkDir creates a directory func (l *localFileSystem) MkDir(path string, mode int) error { fileMode := os.FileMode(mode & int(os.ModePerm)) @@ -93,14 +99,17 @@ func makeFileInfo(filePath string) (FileInfo, error) { } return &localFileInfo{ - filePath: filePath, - fileInfo: f}, nil + filePath: filePath, + fileInfo: f, + accessTime: atime(f), + }, nil } // localFileInfo is implementation of FileInfo for local files type localFileInfo struct { - filePath string - fileInfo os.FileInfo + filePath string + fileInfo os.FileInfo + accessTime time.Time } // IsDir tells this is a directory @@ -152,3 +161,13 @@ func (l *localFileInfo) ReadDir() ([]FileInfo, error) { func (l *localFileInfo) GetModePerm() os.FileMode { return l.fileInfo.Mode() & os.ModePerm } + +// GetModTime returns file modification time +func (l *localFileInfo) GetModTime() time.Time { + return l.fileInfo.ModTime() +} + +// GetAccessTime returns file last access time +func (l *localFileInfo) GetAccessTime() time.Time { + return l.accessTime +} diff --git a/lib/sshutils/scp/scp.go b/lib/sshutils/scp/scp.go index 1ac8bd69499..e6748c5255f 100644 --- a/lib/sshutils/scp/scp.go +++ b/lib/sshutils/scp/scp.go @@ -14,7 +14,12 @@ See the License for the specific language governing permissions and limitations under the License. */ -// Package scp handles file uploads and downloads via scp command +// Package scp handles file uploads and downloads via SCP command. +// See https://web.archive.org/web/20170215184048/https://blogs.oracle.com/janp/entry/how_the_scp_protocol_works +// for the high-level protocol overview. +// +// Authoritative source for the protocol is the source code for OpenSSH scp: +// https://github.com/openssh/openssh-portable/blob/add926dd1bbe3c4db06e27cab8ab0f9a3d00a0c2/scp.c package scp import ( @@ -36,7 +41,7 @@ import ( ) const ( - // OKByte is scp OK message bytes + // OKByte is SCP OK message bytes OKByte = 0x0 // WarnByte tells that next goes a warning string WarnByte = 0x1 @@ -62,6 +67,9 @@ type Flags struct { LocalAddr string // DirectoryMode indicates that a directory is being sent. DirectoryMode bool + // PreserveAttrs preserves access and modification times + // from the original file + PreserveAttrs bool } // Config describes Command configuration settings @@ -105,11 +113,13 @@ type FileSystem interface { OpenFile(filePath string) (io.ReadCloser, error) // CreateFile creates a new file CreateFile(filePath string, length uint64) (io.WriteCloser, error) - // SetChmod sets file permissions - SetChmod(path string, mode int) error + // Chmod sets file permissions + Chmod(path string, mode int) error + // Chtimes sets file access and modification time + Chtimes(path string, atime, mtime time.Time) error } -// FileInfo is an API that describes methods that provide file information +// FileInfo provides access to file metadata type FileInfo interface { // IsDir returns true if a file is a directory IsDir() bool @@ -123,6 +133,10 @@ type FileInfo interface { GetModePerm() os.FileMode // GetSize returns file size GetSize() int64 + // GetModTime returns file modification time + GetModTime() time.Time + // GetAccessTime returns file last access time + GetAccessTime() time.Time } // CreateDownloadCommand configures and returns a command used @@ -164,7 +178,8 @@ func (c *Config) CheckAndSetDefaults() error { return nil } -// CreateCommand creates and returns a new Command +// CreateCommand creates and returns a new SCP command with +// specified configuration. func CreateCommand(cfg Config) (Command, error) { err := cfg.CheckAndSetDefaults() if err != nil { @@ -181,6 +196,7 @@ func CreateCommand(cfg Config) (Command, error) { "LocalAddr": cfg.Flags.LocalAddr, "RemoteAddr": cfg.Flags.RemoteAddr, "Target": cfg.Flags.Target, + "PreserveAttrs": cfg.Flags.PreserveAttrs, "User": cfg.User, "RunOnServer": cfg.RunOnServer, "RemoteLocation": cfg.RemoteLocation, @@ -191,7 +207,7 @@ func CreateCommand(cfg Config) (Command, error) { } // Command mimics behavior of SCP command line tool -// to teleport can pretend it launches real scp behind the scenes +// to teleport can pretend it launches real SCP behind the scenes type command struct { Config log *log.Entry @@ -201,23 +217,22 @@ type command struct { // and teleport (server) side. func (cmd *command) Execute(ch io.ReadWriter) (err error) { if cmd.Flags.Source { - err = cmd.serveSource(ch) - } else { - err = cmd.serveSink(ch) + return trace.Wrap(cmd.serveSource(ch)) } - if err != nil { - return trace.Wrap(err) - } - return nil + return trace.Wrap(cmd.serveSink(ch)) } -func (cmd *command) GetRemoteShellCmd() (string, error) { +// GetRemoteShellCmd returns a command line to copy +// file(s) or a directory to a remote location +func (cmd *command) GetRemoteShellCmd() (shellCmd string, err error) { if cmd.RemoteLocation == "" { return "", trace.BadParameter("missing remote file location") } - // "impersonate" scp to a server - shellCmd := "/usr/bin/scp -f" + // "impersonate" SCP to a server + // See https://docstore.mik.ua/orelly/networking_2ndEd/ssh/ch03_08.htm, section "scp1 Details" + // about the hidden to/from switches + shellCmd = "/usr/bin/scp -f" if cmd.Flags.Source { shellCmd = "/usr/bin/scp -t" } @@ -228,6 +243,9 @@ func (cmd *command) GetRemoteShellCmd() (string, error) { if cmd.Flags.DirectoryMode { shellCmd += " -d" } + if cmd.Flags.PreserveAttrs { + shellCmd += " -p" + } shellCmd += (" " + cmd.RemoteLocation) return shellCmd, nil @@ -248,12 +266,13 @@ func (cmd *command) serveSource(ch io.ReadWriter) (retErr error) { for i := range cmd.Flags.Target { fileInfo, err := cmd.FileSystem.GetFileInfo(cmd.Flags.Target[i]) if err != nil { - err := trace.Errorf("could not access local path %q: %v", cmd.Flags.Target[i], err) - return trace.Wrap(err) + return trace.Errorf("could not access local path %q: %v", cmd.Flags.Target[i], err) } if fileInfo.IsDir() && !cmd.Flags.Recursive { - err := trace.Errorf("%v is a directory, perhaps try -r flag?", fileInfo.GetName()) - return trace.Wrap(err) + // Note: using any other error constructor (e.g. BadParameter) + // might lead to relogin attempt and a completely obscure + // error message + return trace.Errorf("%v is a directory, use -r flag to copy recursively", fileInfo.GetName()) } fileInfos[i] = fileInfo } @@ -276,18 +295,17 @@ func (cmd *command) serveSource(ch io.ReadWriter) (retErr error) { } } - cmd.log.Debugf("send completed") + cmd.log.Debug("Send completed.") return nil } func (cmd *command) sendDir(r *reader, ch io.ReadWriter, fileInfo FileInfo) error { - out := fmt.Sprintf("D%04o 0 %s\n", fileInfo.GetModePerm(), fileInfo.GetName()) - cmd.log.Debugf("sendDir: %v", out) - _, err := io.WriteString(ch, out) - if err != nil { - return trace.Wrap(err) + if cmd.Config.Flags.PreserveAttrs { + if err := cmd.sendFileTimes(r, ch, fileInfo); err != nil { + return trace.Wrap(err) + } } - if err := r.read(); err != nil { + if err := cmd.sendDirMode(r, ch, fileInfo); err != nil { return trace.Wrap(err) } @@ -323,23 +341,15 @@ func (cmd *command) sendFile(r *reader, ch io.ReadWriter, fileInfo FileInfo) err if err != nil { return trace.Wrap(err) } - defer reader.Close() - out := fmt.Sprintf("C%04o %d %s\n", fileInfo.GetModePerm(), fileInfo.GetSize(), fileInfo.GetName()) - - // report progress: - if cmd.ProgressWriter != nil { - statusMessage := fmt.Sprintf("-> %s (%d)", fileInfo.GetPath(), fileInfo.GetSize()) - defer fmt.Fprintf(cmd.ProgressWriter, utils.EscapeControl(statusMessage)+"\n") + if cmd.Config.Flags.PreserveAttrs { + if err := cmd.sendFileTimes(r, ch, fileInfo); err != nil { + return trace.Wrap(err) + } } - _, err = io.WriteString(ch, out) - if err != nil { - return trace.Wrap(err) - } - - if err := r.read(); err != nil { + if err := cmd.sendFileMode(r, ch, fileInfo); err != nil { return trace.Wrap(err) } @@ -348,8 +358,13 @@ func (cmd *command) sendFile(r *reader, ch io.ReadWriter, fileInfo FileInfo) err return trace.Wrap(err) } if n != fileInfo.GetSize() { - err := fmt.Errorf("short write: %v %v", n, fileInfo.GetSize()) - return trace.Wrap(err) + return trace.Errorf("short write: written %v, expected %v", n, fileInfo.GetSize()) + } + + // report progress: + if cmd.ProgressWriter != nil { + statusMessage := fmt.Sprintf("-> %s (%d)", fileInfo.GetPath(), fileInfo.GetSize()) + defer fmt.Fprintf(cmd.ProgressWriter, utils.EscapeControl(statusMessage)+"\n") } if err := sendOK(ch); err != nil { return trace.Wrap(err) @@ -360,18 +375,18 @@ func (cmd *command) sendFile(r *reader, ch io.ReadWriter, fileInfo FileInfo) err func (cmd *command) sendErr(ch io.Writer, err error) { out := fmt.Sprintf("%c%s\n", byte(ErrByte), err) if _, err := ch.Write([]byte(out)); err != nil { - log.Debugf("failed sending SCP error message to the remote side: %v", err) + cmd.log.Debugf("failed sending SCP error message to the remote side: %v", err) } } // serveSink executes file uploading, when a remote server sends file(s) -// via scp +// via SCP func (cmd *command) serveSink(ch io.ReadWriter) error { // Validate that if directory mode flag was sent, the target is an actual // directory. if cmd.Flags.DirectoryMode { if len(cmd.Flags.Target) != 1 { - return trace.BadParameter("in directory mode, only single upload target is allowed, %v provided", len(cmd.Flags.Target)) + return trace.BadParameter("in directory mode, only single upload target is allowed but %v provided", len(cmd.Flags.Target)) } fi, err := os.Stat(cmd.Flags.Target[0]) @@ -387,7 +402,7 @@ func (cmd *command) serveSink(ch io.ReadWriter) error { return trace.Wrap(err) } var st state - st.path = []string{"."} + st.path = localDir var b = make([]byte, 1) scanner := bufio.NewScanner(ch) for { @@ -421,11 +436,9 @@ func (cmd *command) serveSink(ch io.ReadWriter) error { } func (cmd *command) processCommand(ch io.ReadWriter, st *state, b byte, line string) error { - cmd.log.Debugf("[SCP] <- %v %v", string(b), line) + cmd.log.Debugf("<- %v %v", string(b), line) switch b { - case WarnByte: - return trace.Errorf("error from sender: %q", line) - case ErrByte: + case WarnByte, ErrByte: return trace.Errorf("error from sender: %q", line) case 'C': f, err := parseNewFile(line) @@ -447,31 +460,37 @@ func (cmd *command) processCommand(ch io.ReadWriter, st *state, b byte, line str } return nil case 'E': - return st.pop() + if len(st.path) == 0 { + return trace.Errorf("empty path") + } + return cmd.updateDirTimes(st.pop()) case 'T': - _, err := parseMtime(line) + stat, err := parseFileTimes(line) if err != nil { return trace.Wrap(err) } + st.stat = stat + return nil } return trace.Errorf("got unrecognized command: %v", string(b)) } func (cmd *command) receiveFile(st *state, fc newFileCmd, ch io.ReadWriter) error { - cmd.log.Debugf("scp.receiveFile(%v)", cmd.Flags.Target) + cmd.log.Debugf("scp.receiveFile(%v): %v", cmd.Flags.Target, fc.Name) - // if the dest path is a folder, we should save the file to that folder, but - // only if is 'recursive' is set + // if the destination path is a folder, we should save the file to that folder, but + // only if 'recursive' is set path := cmd.Flags.Target[0] if cmd.Flags.Recursive || cmd.FileSystem.IsDir(path) { - path = st.makePath(path, fc.Name) + path = st.makePath(fc.Name) } writer, err := cmd.FileSystem.CreateFile(path, fc.Length) if err != nil { return trace.Wrap(err) } + defer writer.Close() // report progress: if cmd.ProgressWriter != nil { @@ -479,8 +498,6 @@ func (cmd *command) receiveFile(st *state, fc newFileCmd, ch io.ReadWriter) erro defer fmt.Fprintf(cmd.ProgressWriter, utils.EscapeControl(statusMessage)+"\n") } - defer writer.Close() - if err = sendOK(ch); err != nil { return trace.Wrap(err) } @@ -495,31 +512,93 @@ func (cmd *command) receiveFile(st *state, fc newFileCmd, ch io.ReadWriter) erro return trace.Errorf("unexpected file copy length: %v", n) } - if err := cmd.FileSystem.SetChmod(path, int(fc.Mode)); err != nil { + if err := cmd.FileSystem.Chmod(path, int(fc.Mode)); err != nil { return trace.Wrap(err) } + if st.stat != nil { + err = cmd.FileSystem.Chtimes(path, st.stat.Atime, st.stat.Mtime) + if err != nil { + return trace.Wrap(err) + } + } - cmd.log.Debugf("file %v(%v) copied to %v", fc.Name, fc.Length, path) + cmd.log.Debugf("File %v(%v) copied to %v.", fc.Name, fc.Length, path) return nil } func (cmd *command) receiveDir(st *state, fc newFileCmd, ch io.ReadWriter) error { + cmd.log.Debugf("scp.receiveDir(%v): %v", cmd.Flags.Target, fc.Name) targetDir := cmd.Flags.Target[0] // copying into an existing directory? append to it: if cmd.FileSystem.IsDir(targetDir) { - targetDir = st.makePath(targetDir, fc.Name) - st.push(fc.Name) + targetDir = st.makePath(fc.Name) } + st.push(fc.Name, st.stat) err := cmd.FileSystem.MkDir(targetDir, int(fc.Mode)) if err != nil { - return trace.Wrap(err) + return trace.ConvertSystemError(err) } return nil } +func (cmd *command) sendDirMode(r *reader, ch io.Writer, fileInfo FileInfo) error { + out := fmt.Sprintf("D%04o 0 %s\n", fileInfo.GetModePerm(), fileInfo.GetName()) + cmd.log.WithField("cmd", out).Debug("Send directory mode.") + _, err := io.WriteString(ch, out) + if err != nil { + return trace.Wrap(err) + } + return trace.Wrap(r.read()) +} + +func (cmd *command) sendFileTimes(r *reader, ch io.Writer, fileInfo FileInfo) error { + // OpenSSH handles nanoseconds to a certain precision + // which is not sufficient to keep the exact timestamps: + // See these for details: + // https://github.com/openssh/openssh-portable/blob/279261e1ea8150c7c64ab5fe7cb4a4ea17acbb29/scp.c#L619-L621 + // https://github.com/openssh/openssh-portable/blob/279261e1ea8150c7c64ab5fe7cb4a4ea17acbb29/scp.c#L1332 + // https://github.com/openssh/openssh-portable/blob/279261e1ea8150c7c64ab5fe7cb4a4ea17acbb29/scp.c#L1344 + // + // Se we copy its behavior and drop nanoseconds entirely + out := fmt.Sprintf("T%d 0 %d 0\n", + fileInfo.GetModTime().Unix(), + fileInfo.GetAccessTime().Unix(), + ) + cmd.log.WithField("cmd", out).Debug("Send file times.") + _, err := io.WriteString(ch, out) + if err != nil { + return trace.Wrap(err) + } + return trace.Wrap(r.read()) +} + +func (cmd *command) sendFileMode(r *reader, ch io.Writer, fileInfo FileInfo) error { + out := fmt.Sprintf("C%04o %d %s\n", + fileInfo.GetModePerm(), + fileInfo.GetSize(), + fileInfo.GetName(), + ) + cmd.log.WithField("cmd", out).Debug("Send file mode.") + _, err := io.WriteString(ch, out) + if err != nil { + return trace.Wrap(err) + } + return trace.Wrap(r.read()) +} + +func (cmd *command) updateDirTimes(path pathSegments) error { + if stat := path[len(path)-1].stat; stat != nil { + err := cmd.FileSystem.Chtimes(path.join(), stat.Atime, stat.Mtime) + if err != nil { + return trace.ConvertSystemError(err) + } + } + return nil +} + type newFileCmd struct { Mode int64 Length uint64 @@ -558,7 +637,13 @@ type mtimeCmd struct { Atime time.Time } -func parseMtime(line string) (*mtimeCmd, error) { +// parseFileTimes parses the input with access/modification file times: +// +// T +// +// Note that the leading 'T' will not be part of the input as it has already +// been seen and removed +func parseFileTimes(line string) (*mtimeCmd, error) { parts := strings.SplitN(line, " ", 4) if len(parts) != 4 { return nil, trace.Errorf("broken mtime command") @@ -585,28 +670,47 @@ func sendOK(ch io.ReadWriter) error { } type state struct { - path []string - finished bool + path pathSegments + // stat optionally specifies access/modification time for the current file/directory + stat *mtimeCmd } -func (st *state) push(dir string) { - st.path = append(st.path, dir) -} - -func (st *state) pop() error { - if st.finished { - return trace.Errorf("empty path") +func (r pathSegments) join() string { + path := make([]string, 0, len(r)) + for _, s := range r { + path = append(path, s.dir) } + return filepath.Join(path...) +} + +var localDir = pathSegments{{dir: "."}} + +type pathSegments []pathSegment + +type pathSegment struct { + dir string + // stat optionally specifies access/modification time for the directory + stat *mtimeCmd +} + +func (st *state) push(dir string, stat *mtimeCmd) { + st.path = append(st.path, pathSegment{dir: dir, stat: stat}) +} + +// pop removes the last segment from the current path. +// Returns the old path as a result +func (st *state) pop() pathSegments { if len(st.path) == 0 { - st.finished = true // allow extra 'E' command in the end return nil } + path := st.path st.path = st.path[:len(st.path)-1] - return nil + st.stat = nil + return path } -func (st *state) makePath(target, filename string) string { - return filepath.Join(target, filepath.Join(st.path...), filename) +func (st *state) makePath(filename string) string { + return filepath.Join(st.path.join(), filename) } func newReader(r io.Reader) *reader { @@ -664,31 +768,41 @@ var reSCP = regexp.MustCompile( `(?:[^@\[\:\]]+)` + `)` + // after colon, there is a path that could consist technically of - // any char - `:(?P.+)`, + // any char including empty which stands for the implicit home directory + `:(?P.*)`, ) -// Destination is scp destination to copy to or from +// Destination is SCP destination to copy to or from type Destination struct { // Login is an optional login username Login string // Host is a host to copy to/from Host utils.NetAddr - // Path is a path to copy to/from + // Path is a path to copy to/from. + // An empty path name is valid, and it refers to the user's default directory (usually + // the user's home directory). + // See https://tools.ietf.org/html/draft-ietf-secsh-filexfer-09#page-14, 'File Names' Path string } // ParseSCPDestination takes a string representing a remote resource for SCP -// to download/upload, like "user@host:/path/to/resource.txt" and returns -// 3 components of it +// to download/upload, like "user@host:/path/to/resource.txt" and parses it into +// a structured form. +// +// See https://tools.ietf.org/html/draft-ietf-secsh-filexfer-09#page-14, 'File Names' +// section about details on file names. func ParseSCPDestination(s string) (*Destination, error) { out := reSCP.FindStringSubmatch(s) - if len(out) == 0 { + if len(out) < 4 { return nil, trace.BadParameter("failed to parse %q, try form user@host:/path", s) } addr, err := utils.ParseAddr(out[2]) if err != nil { return nil, trace.Wrap(err) } - return &Destination{Login: out[1], Host: *addr, Path: out[3]}, nil + path := out[3] + if path == "" { + path = "." + } + return &Destination{Login: out[1], Host: *addr, Path: path}, nil } diff --git a/lib/sshutils/scp/scp_test.go b/lib/sshutils/scp/scp_test.go index c1ddb81b705..03147f82d1d 100644 --- a/lib/sshutils/scp/scp_test.go +++ b/lib/sshutils/scp/scp_test.go @@ -1,5 +1,5 @@ /* -Copyright 2018 Gravitational, Inc. +Copyright 2018-2020 Gravitational, Inc. Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the License. @@ -31,27 +31,19 @@ import ( "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" - . "gopkg.in/check.v1" + + "github.com/google/go-cmp/cmp" + "github.com/sirupsen/logrus" + "github.com/stretchr/testify/require" ) -func TestSCP(t *testing.T) { TestingT(t) } +func TestHTTPSendFile(t *testing.T) { + outDir := tempDir(t) -type SCPSuite struct { -} - -var _ = fmt.Printf -var _ = Suite(&SCPSuite{}) - -func (s *SCPSuite) SetUpSuite(c *C) { - utils.InitLoggerForTests(testing.Verbose()) -} - -func (s *SCPSuite) TestHTTPSendFile(c *C) { - outdir := c.MkDir() expectedBytes := []byte("hello") buf := bytes.NewReader(expectedBytes) req, err := http.NewRequest("POST", "/", buf) - c.Assert(err, IsNil) + require.NoError(t, err) req.Header.Set("Content-Length", strconv.Itoa(len(expectedBytes))) @@ -59,26 +51,25 @@ func (s *SCPSuite) TestHTTPSendFile(c *C) { cmd, err := CreateHTTPUpload( HTTPTransferRequest{ FileName: "filename", - RemoteLocation: outdir, + RemoteLocation: outDir, HTTPRequest: req, Progress: stdOut, User: "test-user", }) - c.Assert(err, IsNil) - err = runSCP(cmd, "scp", "-v", "-t", outdir) - c.Assert(err, IsNil) - bytesReceived, err := ioutil.ReadFile(filepath.Join(outdir, "filename")) - c.Assert(err, IsNil) - c.Assert(string(bytesReceived), Equals, string(expectedBytes)) + require.NoError(t, err) + err = runSCP(cmd, "-v", "-t", outDir) + require.NoError(t, err) + bytesReceived, err := ioutil.ReadFile(filepath.Join(outDir, "filename")) + require.NoError(t, err) + require.Empty(t, cmp.Diff(string(bytesReceived), string(expectedBytes))) } -func (s *SCPSuite) TestHTTPReceiveFile(c *C) { - dir := c.MkDir() - source := filepath.Join(dir, "target") +func TestHTTPReceiveFile(t *testing.T) { + source := filepath.Join(tempDir(t), "target") contents := []byte("hello, file contents!") err := ioutil.WriteFile(source, contents, 0666) - c.Assert(err, IsNil) + require.NoError(t, err) w := httptest.NewRecorder() stdOut := bytes.NewBufferString("") @@ -89,182 +80,156 @@ func (s *SCPSuite) TestHTTPReceiveFile(c *C) { User: "test-user", Progress: stdOut, }) + require.NoError(t, err) - c.Assert(err, IsNil) - - err = runSCP(cmd, "scp", "-v", "-f", source) - c.Assert(err, IsNil) + err = runSCP(cmd, "-v", "-f", source) + require.NoError(t, err) data, err := ioutil.ReadAll(w.Body) contentLengthStr := strconv.Itoa(len(data)) - c.Assert(err, IsNil) - c.Assert(string(data), Equals, string(contents)) - c.Assert(contentLengthStr, Equals, w.Header().Get("Content-Length")) - c.Assert("application/octet-stream", Equals, w.Header().Get("Content-Type")) - c.Assert(`attachment;filename="robots.txt"`, Equals, w.Header().Get("Content-Disposition")) + require.NoError(t, err) + require.Empty(t, cmp.Diff(string(data), string(contents))) + require.Empty(t, cmp.Diff(contentLengthStr, w.Header().Get("Content-Length"))) + require.Empty(t, cmp.Diff("application/octet-stream", w.Header().Get("Content-Type"))) + require.Empty(t, cmp.Diff(`attachment;filename="robots.txt"`, w.Header().Get("Content-Disposition"))) } -func (s *SCPSuite) TestSendFile(c *C) { - dir := c.MkDir() - target := filepath.Join(dir, "target") - - cmd, err := CreateCommand( - Config{ - User: "test-user", - Flags: Flags{ - Source: true, - Target: []string{target}, - }, +func TestSend(t *testing.T) { + t.Parallel() + utils.InitLoggerForTests(testing.Verbose()) + modtime := testNow + atime := testNow.Add(1 * time.Second) + dirModtime := testNow.Add(2 * time.Second) + dirAtime := testNow.Add(3 * time.Second) + logger := logrus.WithField(trace.Component, "t:send") + var testCases = []struct { + desc string + config Config + fs testFS + args []string + }{ + { + desc: "regular file preserving the attributes", + config: newSourceConfig("file", Flags{PreserveAttrs: true}), + args: args("-v", "-t", "-p"), + fs: newTestFS(logger, newFile("file", modtime, atime, "file contents")), }, - ) - c.Assert(err, IsNil) - - // Source file is missing, expect an error. - outDir := c.MkDir() - err = runSCP(cmd, "scp", "-v", "-t", outDir) - c.Assert(err, NotNil) - - _, err = ioutil.ReadFile(filepath.Join(outDir, "target")) - c.Assert(os.IsNotExist(err), Equals, true) - - // Write the source file. - contents := []byte("hello, send file!") - err = ioutil.WriteFile(target, contents, 0666) - c.Assert(err, IsNil) - - // Source file is present, should send fine. - err = runSCP(cmd, "scp", "-v", "-t", outDir) - c.Assert(err, IsNil) - - bytes, err := ioutil.ReadFile(filepath.Join(outDir, "target")) - c.Assert(err, IsNil) - c.Assert(string(bytes), Equals, string(contents)) -} - -func (s *SCPSuite) TestReceiveFile(c *C) { - dir := c.MkDir() - source := filepath.Join(dir, "target") - - outDir := c.MkDir() + "/" - cmd, err := CreateCommand(Config{ - User: "test-user", - Flags: Flags{ - Sink: true, - Target: []string{outDir}, + { + desc: "directory preserving the attributes", + config: newSourceConfig("dir", Flags{PreserveAttrs: true, Recursive: true}), + args: args("-v", "-t", "-r", "-p"), + fs: newTestFS( + logger, + // Use timestamps extending backwards to test time application + newDir("dir", dirModtime.Add(1*time.Second), dirAtime.Add(2*time.Second), + newFile("dir/file", modtime.Add(1*time.Minute), atime.Add(2*time.Minute), "file contents"), + newDir("dir/dir2", dirModtime, dirAtime, + newFile("dir/dir2/file2", modtime, atime, "file2 contents")), + ), + ), }, - }) - c.Assert(err, IsNil) + } + for _, tt := range testCases { + tt := tt + t.Run(tt.desc, func(t *testing.T) { + t.Parallel() + cmd, err := CreateCommand(tt.config) + require.NoError(t, err) - // Source file is missing, expect an error. - err = runSCP(cmd, "scp", "-v", "-f", source) - c.Assert(err, NotNil) + targetDir := tempDir(t) + target := filepath.Join(targetDir, tt.config.Flags.Target[0]) + args := append(tt.args, target) - _, err = ioutil.ReadFile(filepath.Join(outDir, "target")) - c.Assert(os.IsNotExist(err), Equals, true) + // Source is missing, expect an error. + err = runSCP(cmd, args...) + require.Regexp(t, "could not access local path.*no such file or directory", err) - // Write the source file. - contents := []byte("hello, file contents!") - err = ioutil.WriteFile(source, contents, 0666) - c.Assert(err, IsNil) + tt.config.FileSystem = tt.fs + cmd, err = CreateCommand(tt.config) + require.NoError(t, err) + // Resend the data + err = runSCP(cmd, args...) + require.NoError(t, err) - // Source file is present, should send fine. - err = runSCP(cmd, "scp", "-v", "-f", source) - c.Assert(err, IsNil) - - bytes, err := ioutil.ReadFile(filepath.Join(outDir, "target")) - c.Assert(err, IsNil) - c.Assert(string(bytes), Equals, string(contents)) + fs := newEmptyTestFS(logger) + fromOS(t, targetDir, &fs) + validateSCP(t, fs, tt.fs) + validateSCPContents(t, fs, tt.fs) + }) + } } -func (s *SCPSuite) TestSendDir(c *C) { - dir := filepath.Join(c.MkDir(), "target_dir") - - cmd, err := CreateCommand(Config{ - User: "test-user", - Flags: Flags{ - Source: true, - Target: []string{dir}, - Recursive: true, +func TestReceive(t *testing.T) { + t.Parallel() + utils.InitLoggerForTests(testing.Verbose()) + modtime := testNow + atime := testNow.Add(1 * time.Second) + dirModtime := testNow.Add(2 * time.Second) + dirAtime := testNow.Add(3 * time.Second) + logger := logrus.WithField(trace.Component, "t:recv") + var testCases = []struct { + desc string + config Config + fs testFS + args []string + }{ + { + desc: "regular file preserving the attributes", + config: newTargetConfig("file", Flags{PreserveAttrs: true}), + args: args("-v", "-f", "-p"), + fs: newTestFS(logger, newFile("file", modtime, atime, "file contents")), }, - }) - c.Assert(err, IsNil) - - // Source directory is missing, expect an error. - outDir := c.MkDir() - err = runSCP(cmd, "scp", "-v", "-r", "-t", outDir) - c.Assert(err, NotNil) - _, err = os.Stat(filepath.Join(outDir, filepath.Base(dir))) - c.Assert(os.IsNotExist(err), Equals, true) - - // Create an populate source directory. - c.Assert(os.MkdirAll(filepath.Join(dir, "nested_dir"), 0777), IsNil) - - err = ioutil.WriteFile( - filepath.Join(dir, "nested_dir", "target1"), []byte("file 1"), 0666) - c.Assert(err, IsNil) - - err = ioutil.WriteFile( - filepath.Join(dir, "target2"), []byte("file 2"), 0666) - c.Assert(err, IsNil) - - // Source directory is present, should send fine. - err = runSCP(cmd, "scp", "-v", "-r", "-t", outDir) - c.Assert(err, IsNil) - - name := filepath.Base(dir) - bytes, err := ioutil.ReadFile(filepath.Join(outDir, name, "nested_dir", "target1")) - c.Assert(err, IsNil) - c.Assert(string(bytes), Equals, string("file 1")) - - bytes, err = ioutil.ReadFile(filepath.Join(outDir, name, "target2")) - c.Assert(err, IsNil) - c.Assert(string(bytes), Equals, string("file 2")) -} - -func (s *SCPSuite) TestReceiveDir(c *C) { - dir := filepath.Join(c.MkDir(), "target_dir") - - outDir := c.MkDir() + "/" - cmd, err := CreateCommand(Config{ - User: "test-user", - Flags: Flags{ - Sink: true, - Target: []string{outDir}, - Recursive: true, + { + desc: "directory preserving the attributes", + config: newTargetConfig("dir", Flags{PreserveAttrs: true, Recursive: true}), + args: args("-v", "-f", "-r", "-p"), + fs: newTestFS( + logger, + // Use timestamps extending backwards to test time application + newDir("dir", dirModtime.Add(1*time.Second), dirAtime.Add(2*time.Second), + newFile("dir/file", modtime.Add(1*time.Minute), atime.Add(2*time.Minute), "file contents"), + newDir("dir/dir2", dirModtime, dirAtime, + newFile("dir/dir2/file2", modtime, atime, "file2 contents")), + ), + ), }, - }) - c.Assert(err, IsNil) + } + for _, tt := range testCases { + tt := tt + t.Run(tt.desc, func(t *testing.T) { + t.Parallel() + cmd, err := CreateCommand(tt.config) + require.NoError(t, err) - // Source directory is missing, expect an error. - err = runSCP(cmd, "scp", "-v", "-r", "-f", dir) - c.Assert(err, NotNil) + sourceDir := tempDir(t) + source := filepath.Join(sourceDir, tt.config.Flags.Target[0]) + args := append(tt.args, source) - // Create an populate source directory. - c.Assert(os.MkdirAll(filepath.Join(dir, "nested_dir"), 0777), IsNil) + // Source is missing, expect an error. + err = runSCP(cmd, args...) + require.Regexp(t, ".*No such file or directory", err) - err = ioutil.WriteFile( - filepath.Join(dir, "nested_dir", "target1"), []byte("file 1"), 0666) - c.Assert(err, IsNil) + fs := newEmptyTestFS(logger) + tt.config.FileSystem = fs + cmd, err = CreateCommand(tt.config) + require.NoError(t, err) - err = ioutil.WriteFile( - filepath.Join(dir, "target2"), []byte("file 2"), 0666) - c.Assert(err, IsNil) + writeData(t, sourceDir, tt.fs) + writeFileTimes(t, sourceDir, tt.fs) - // Source directory is present, should send fine. - err = runSCP(cmd, "scp", "-v", "-r", "-f", dir) - c.Assert(err, IsNil) + // Resend the data + err = runSCP(cmd, args...) + require.NoError(t, err) - name := filepath.Base(dir) - bytes, err := ioutil.ReadFile(filepath.Join(outDir, name, "nested_dir", "target1")) - c.Assert(err, IsNil) - c.Assert(string(bytes), Equals, string("file 1")) - - bytes, err = ioutil.ReadFile(filepath.Join(outDir, name, "target2")) - c.Assert(err, IsNil) - c.Assert(string(bytes), Equals, string("file 2")) + validateSCP(t, tt.fs, fs) + validateSCPContents(t, tt.fs, fs) + }) + } } -func (s *SCPSuite) TestInvalidDir(c *C) { +func TestInvalidDir(t *testing.T) { + t.Parallel() + cmd, err := CreateCommand(Config{ User: "test-user", Flags: Flags{ @@ -273,37 +238,54 @@ func (s *SCPSuite) TestInvalidDir(c *C) { Recursive: true, }, }) - c.Assert(err, IsNil) + require.NoError(t, err) - tests := []struct { + testCases := []struct { + desc string inDirName string + err string }{ - {inDirName: ""}, - {inDirName: "."}, - {inDirName: ".."}, + { + desc: "no directory", + inDirName: "", + err: ".*No such file or directory.*", + }, + { + desc: "current directory", + inDirName: ".", + err: ".*invalid name.*", + }, + { + desc: "parent directory", + inDirName: "..", + err: ".*invalid name.*", + }, } - for _, tt := range tests { - scp, in, out, _ := run("scp", "-v", "-r", "-f", tt.inDirName) - rw := &combo{out, in} + for _, tt := range testCases { + tt := tt + t.Run(tt.desc, func(t *testing.T) { + scp, in, out, _ := newCmd("scp", "-v", "-r", "-f", tt.inDirName) + rw := &readWriter{out, in} - err := scp.Start() - c.Assert(err, IsNil) + err := scp.Start() + require.NoError(t, err) - err = cmd.Execute(rw) - c.Assert(err, NotNil) + err = cmd.Execute(rw) + require.Regexp(t, tt.err, err) + }) } } // TestVerifyDir makes sure that if scp was started in directory mode (the // user attempts to copy multiple files or a directory), the target is a // directory. -func (s *SCPSuite) TestVerifyDir(c *C) { +func TestVerifyDir(t *testing.T) { // Create temporary directory with a file "target" in it. - dir := c.MkDir() + dir := tempDir(t) target := filepath.Join(dir, "target") err := ioutil.WriteFile(target, []byte{}, 0666) - c.Assert(err, IsNil) + require.NoError(t, err) cmd, err := CreateCommand( Config{ @@ -314,63 +296,77 @@ func (s *SCPSuite) TestVerifyDir(c *C) { }, }, ) - c.Assert(err, IsNil) + require.NoError(t, err) // Run command with -d flag (directory mode). Since the target is a file, // it should fail. - err = runSCP(cmd, "scp", "-t", "-d", target) - c.Assert(err, NotNil) + err = runSCP(cmd, "-t", "-d", target) + require.Regexp(t, ".*Not a directory", err) } -func (s *SCPSuite) TestSCPParsing(c *C) { - type tc struct { - in string - dest Destination - err error - } - testCases := []tc{ +func TestSCPParsing(t *testing.T) { + t.Parallel() + + var testCases = []struct { + comment string + in string + dest Destination + err error + }{ { - in: "root@remote.host:/etc/nginx.conf", - dest: Destination{Login: "root", Host: utils.NetAddr{Addr: "remote.host", AddrNetwork: "tcp"}, Path: "/etc/nginx.conf"}, + comment: "full spec of the remote destination", + in: "root@remote.host:/etc/nginx.conf", + dest: Destination{Login: "root", Host: utils.NetAddr{Addr: "remote.host", AddrNetwork: "tcp"}, Path: "/etc/nginx.conf"}, }, { - in: "remote.host:/etc/nginx.co:nf", - dest: Destination{Host: utils.NetAddr{Addr: "remote.host", AddrNetwork: "tcp"}, Path: "/etc/nginx.co:nf"}, + comment: "spec with just the remote host", + in: "remote.host:/etc/nginx.co:nf", + dest: Destination{Host: utils.NetAddr{Addr: "remote.host", AddrNetwork: "tcp"}, Path: "/etc/nginx.co:nf"}, }, { - in: "[::1]:/etc/nginx.co:nf", - dest: Destination{Host: utils.NetAddr{Addr: "[::1]", AddrNetwork: "tcp"}, Path: "/etc/nginx.co:nf"}, + comment: "ipv6 remote destination address", + in: "[::1]:/etc/nginx.co:nf", + dest: Destination{Host: utils.NetAddr{Addr: "[::1]", AddrNetwork: "tcp"}, Path: "/etc/nginx.co:nf"}, }, { - in: "root@123.123.123.123:/var/www/html/", - dest: Destination{Login: "root", Host: utils.NetAddr{Addr: "123.123.123.123", AddrNetwork: "tcp"}, Path: "/var/www/html/"}, + comment: "full spec of the remote destination using ipv4 address", + in: "root@123.123.123.123:/var/www/html/", + dest: Destination{Login: "root", Host: utils.NetAddr{Addr: "123.123.123.123", AddrNetwork: "tcp"}, Path: "/var/www/html/"}, }, { - in: "myusername@myremotehost.com:/home/hope/*", - dest: Destination{Login: "myusername", Host: utils.NetAddr{Addr: "myremotehost.com", AddrNetwork: "tcp"}, Path: "/home/hope/*"}, + comment: "target location using wildcard", + in: "myusername@myremotehost.com:/home/hope/*", + dest: Destination{Login: "myusername", Host: utils.NetAddr{Addr: "myremotehost.com", AddrNetwork: "tcp"}, Path: "/home/hope/*"}, }, { - in: "complex@example.com@remote.com:/anything.txt", - dest: Destination{Login: "complex@example.com", Host: utils.NetAddr{Addr: "remote.com", AddrNetwork: "tcp"}, Path: "/anything.txt"}, + comment: "complex login", + in: "complex@example.com@remote.com:/anything.txt", + dest: Destination{Login: "complex@example.com", Host: utils.NetAddr{Addr: "remote.com", AddrNetwork: "tcp"}, Path: "/anything.txt"}, + }, + { + comment: "implicit user's home directory", + in: "root@remote.host:", + dest: Destination{Login: "root", Host: utils.NetAddr{Addr: "remote.host", AddrNetwork: "tcp"}, Path: "."}, }, } - for i, tc := range testCases { - comment := Commentf("Test case %v: %q", i, tc.in) - re, err := ParseSCPDestination(tc.in) - if tc.err == nil { - c.Assert(err, IsNil, comment) - c.Assert(re.Login, Equals, tc.dest.Login, comment) - c.Assert(re.Host, DeepEquals, tc.dest.Host, comment) - c.Assert(re.Path, Equals, tc.dest.Path, comment) - } else { - c.Assert(err, FitsTypeOf, tc.err) - } + for _, tt := range testCases { + tt := tt + t.Run(tt.comment, func(t *testing.T) { + resp, err := ParseSCPDestination(tt.in) + if tt.err != nil { + require.IsType(t, err, tt.err) + return + } + require.NoError(t, err) + require.Empty(t, cmp.Diff(resp, &tt.dest)) + }) + } } -func runSCP(cmd Command, name string, args ...string) error { - scp, in, out, _ := run(name, args...) - rw := &combo{out, in} +func runSCP(cmd Command, args ...string) error { + scp, in, out, _ := newCmd("scp", args...) + rw := &readWriter{out, in} errCh := make(chan error, 1) @@ -402,36 +398,344 @@ func runSCP(cmd Command, name string, args ...string) error { } } -type combo struct { +// fromOS recreates the structure of the specified directory dir +// into the provided file system fs +func fromOS(t *testing.T, dir string, fs *testFS) { + err := filepath.Walk(dir, func(path string, fi os.FileInfo, err error) error { + if err != nil { + return err + } + relpath, err := filepath.Rel(dir, path) + require.NoError(t, err) + if relpath == "." { + // Skip top-level directory + return nil + } + if fi.IsDir() { + require.NoError(t, fs.MkDir(relpath, int(fi.Mode()))) + require.NoError(t, fs.Chtimes(relpath, atime(fi), fi.ModTime())) + return nil + } + wc, err := fs.CreateFile(relpath, uint64(fi.Size())) + require.NoError(t, err) + defer wc.Close() + require.NoError(t, fs.Chtimes(relpath, atime(fi), fi.ModTime())) + f, err := os.Open(path) + require.NoError(t, err) + defer f.Close() + _, err = io.Copy(wc, f) + require.NoError(t, err) + return nil + }) + require.NoError(t, err) +} + +// writeData recreates the file/directory structure in dir +// as specified with the file system fs +func writeData(t *testing.T, dir string, fs testFS) { + for _, f := range fs.fs { + if f.IsDir() { + require.NoError(t, os.MkdirAll(filepath.Join(dir, f.path), f.perms)) + continue + } + rc, err := fs.OpenFile(f.path) + require.NoError(t, err) + defer rc.Close() + targetPath := filepath.Join(dir, f.path) + if parentDir := filepath.Dir(f.path); parentDir != "." { + fi := fs.fs[parentDir] + require.NoError(t, os.MkdirAll(filepath.Dir(targetPath), fi.perms)) + } + f, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY, f.perms) + require.NoError(t, err) + defer f.Close() + _, err = io.Copy(f, rc) + require.NoError(t, err) + } +} + +// writeFileTimes applies access/modification times on files/directories in dir +// as specified in the file system fs. +func writeFileTimes(t *testing.T, dir string, fs testFS) { + for _, f := range fs.fs { + require.NoError(t, os.Chtimes(filepath.Join(dir, f.path), f.atime, f.modtime)) + } +} + +// validateSCPContents verifies that the file contents in the specified +// file systems match in the corresponding files +func validateSCPContents(t *testing.T, expected testFS, actual FileSystem) { + for path, fileinfo := range expected.fs { + if fileinfo.IsDir() { + continue + } + rc, err := actual.OpenFile(path) + require.NoError(t, err) + defer rc.Close() + bytes, err := ioutil.ReadAll(rc) + require.NoError(t, err) + require.Empty(t, cmp.Diff(fileinfo.contents.String(), string(bytes))) + } +} + +// validateSCP verifies that the specified pair of FileSystems match. +// FileSystem match if their contents match incl. access/modification times +func validateSCP(t *testing.T, expected testFS, actual FileSystem) { + for path, fileinfo := range expected.fs { + targetFileinfo, err := actual.GetFileInfo(path) + require.NoError(t, err) + if fileinfo.IsDir() { + require.True(t, targetFileinfo.IsDir()) + } else { + require.True(t, targetFileinfo.GetModePerm().IsRegular()) + } + validateFileTimes(t, *fileinfo, targetFileinfo) + } +} + +// validateFileTimes verifies that the specified pair of FileInfos match +func validateFileTimes(t *testing.T, expected testFileInfo, actual FileInfo) { + require.Empty(t, cmp.Diff( + expected.GetModTime().UTC().Format(time.RFC3339), + actual.GetModTime().UTC().Format(time.RFC3339), + )) + require.Empty(t, cmp.Diff( + expected.GetAccessTime().UTC().Format(time.RFC3339), + actual.GetAccessTime().UTC().Format(time.RFC3339), + )) +} + +type readWriter struct { r io.Reader w io.Writer } -func (c *combo) Read(b []byte) (int, error) { +func (c *readWriter) Read(b []byte) (int, error) { return c.r.Read(b) } -func (c *combo) Write(b []byte) (int, error) { +func (c *readWriter) Write(b []byte) (int, error) { return c.w.Write(b) } -func run(name string, args ...string) (*exec.Cmd, io.WriteCloser, io.ReadCloser, io.ReadCloser) { - cmd := exec.Command(name, args...) +func newCmd(name string, args ...string) (cmd *exec.Cmd, stdin io.WriteCloser, stdout io.ReadCloser, stderr io.ReadCloser) { + cmd = exec.Command(name, args...) - in, err := cmd.StdinPipe() + var err error + stdin, err = cmd.StdinPipe() if err != nil { panic(err) } - out, err := cmd.StdoutPipe() + stdout, err = cmd.StdoutPipe() if err != nil { panic(err) } - epipe, err := cmd.StderrPipe() + stderr, err = cmd.StderrPipe() if err != nil { panic(err) } - return cmd, in, out, epipe + return cmd, stdin, stdout, stderr +} + +// newEmptyTestFS creates a new test FileSystem without content +func newEmptyTestFS(l logrus.FieldLogger) testFS { + return testFS{ + fs: make(map[string]*testFileInfo), + l: l, + } +} + +// newTestFS creates a new test FileSystem using the specified logger +// and the set of top-level files +func newTestFS(l logrus.FieldLogger, files ...*testFileInfo) testFS { + fs := make(map[string]*testFileInfo) + addFiles(fs, files...) + return testFS{ + fs: fs, + l: l, + } +} + +func (r testFS) IsDir(path string) bool { + r.l.WithField("path", path).Info("IsDir.") + if fi, exists := r.fs[path]; exists { + return fi.IsDir() + } + return false +} + +func (r testFS) GetFileInfo(path string) (FileInfo, error) { + r.l.WithField("path", path).Info("GetFileInfo.") + fi, exists := r.fs[path] + if !exists { + return nil, errMissingFile + } + return fi, nil +} + +func (r testFS) MkDir(path string, mode int) error { + r.l.WithField("path", path).WithField("mode", mode).Info("MkDir.") + _, exists := r.fs[path] + if exists { + return trace.AlreadyExists("directory %v already exists", path) + } + r.fs[path] = &testFileInfo{ + path: path, + dir: true, + perms: os.FileMode(mode) | os.ModeDir, + } + return nil +} + +func (r testFS) OpenFile(path string) (io.ReadCloser, error) { + r.l.WithField("path", path).Info("OpenFile.") + fi, exists := r.fs[path] + if !exists { + return nil, errMissingFile + } + rc := nopReadCloser{Reader: bytes.NewReader(fi.contents.Bytes())} + return rc, nil +} + +func (r testFS) CreateFile(path string, length uint64) (io.WriteCloser, error) { + r.l.WithField("path", path).WithField("len", length).Info("CreateFile.") + fi := &testFileInfo{ + path: path, + size: int64(length), + perms: 0666, + contents: new(bytes.Buffer), + } + r.fs[path] = fi + if dir := filepath.Dir(path); dir != "." { + r.MkDir(dir, 0755) + r.fs[dir].ents = append(r.fs[dir].ents, fi) + } + wc := utils.NopWriteCloser(fi.contents) + return wc, nil +} + +func (r testFS) Chmod(path string, mode int) error { + r.l.WithField("path", path).WithField("mode", mode).Info("Chmod.") + fi, exists := r.fs[path] + if !exists { + return errMissingFile + } + fi.perms = os.FileMode(mode) + return nil +} + +func (r testFS) Chtimes(path string, atime, mtime time.Time) error { + r.l.WithField("path", path).WithField("atime", atime).WithField("mtime", mtime).Info("Chtimes.") + fi, exists := r.fs[path] + if !exists { + return errMissingFile + } + fi.modtime = mtime + fi.atime = atime + return nil +} + +// testFS implements a fake FileSystem +type testFS struct { + l logrus.FieldLogger + fs map[string]*testFileInfo +} + +type testFileInfo struct { + dir bool + perms os.FileMode + path string + modtime time.Time + atime time.Time + ents []*testFileInfo + size int64 + contents *bytes.Buffer +} + +func (r *testFileInfo) IsDir() bool { return r.dir } +func (r *testFileInfo) ReadDir() (fis []FileInfo, err error) { + fis = make([]FileInfo, 0, len(r.ents)) + for _, e := range r.ents { + fis = append(fis, e) + } + return fis, nil +} +func (r *testFileInfo) GetName() string { return filepath.Base(r.path) } +func (r *testFileInfo) GetPath() string { return r.path } +func (r *testFileInfo) GetModePerm() os.FileMode { return r.perms } +func (r *testFileInfo) GetSize() int64 { return r.size } +func (r *testFileInfo) GetModTime() time.Time { return r.modtime } +func (r *testFileInfo) GetAccessTime() time.Time { return r.atime } + +func (r nopReadCloser) Close() error { return nil } + +type nopReadCloser struct { + io.Reader +} + +var errMissingFile = fmt.Errorf("no such file or directory") + +func tempDir(t *testing.T) (dir string) { + path, err := ioutil.TempDir("", "test") + require.NoError(t, err) + t.Cleanup(func() { os.RemoveAll(path) }) + return path +} + +func newSourceConfig(path string, flags Flags) Config { + flags.Source = true + flags.Target = []string{path} + return Config{ + User: "test-user", + Flags: flags, + } +} + +func newTargetConfig(path string, flags Flags) Config { + flags.Sink = true + flags.Target = []string{path} + return Config{ + User: "test-user", + Flags: flags, + } +} + +func newDir(name string, modtime, atime time.Time, ents ...*testFileInfo) *testFileInfo { + return &testFileInfo{ + path: name, + ents: ents, + modtime: modtime, + atime: atime, + dir: true, + perms: 0755, + } +} + +func newFile(name string, modtime, atime time.Time, contents string) *testFileInfo { + return &testFileInfo{ + path: name, + modtime: modtime, + atime: atime, + perms: 0666, + size: int64(len(contents)), + contents: bytes.NewBufferString(contents), + } +} + +func addFiles(fs map[string]*testFileInfo, ents ...*testFileInfo) { + for _, f := range ents { + fs[f.path] = f + if f.IsDir() { + addFiles(fs, f.ents...) + } + } +} + +var testNow = time.Date(1984, time.April, 4, 0, 0, 0, 0, time.UTC) + +func args(params ...string) []string { + return params } diff --git a/lib/sshutils/scp/stat_darwin.go b/lib/sshutils/scp/stat_darwin.go new file mode 100644 index 00000000000..935b691b6d9 --- /dev/null +++ b/lib/sshutils/scp/stat_darwin.go @@ -0,0 +1,32 @@ +/* +Copyright 2020 Gravitational, Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package scp + +import ( + "os" + "syscall" + "time" +) + +// Source: os/stat_darwin.go +func atime(fi os.FileInfo) time.Time { + return timespecToTime(fi.Sys().(*syscall.Stat_t).Atimespec) +} + +func timespecToTime(ts syscall.Timespec) time.Time { + return time.Unix(int64(ts.Sec), int64(ts.Nsec)) +} diff --git a/lib/sshutils/scp/stat_linux.go b/lib/sshutils/scp/stat_linux.go new file mode 100644 index 00000000000..023e080d315 --- /dev/null +++ b/lib/sshutils/scp/stat_linux.go @@ -0,0 +1,32 @@ +/* +Copyright 2020 Gravitational, Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package scp + +import ( + "os" + "syscall" + "time" +) + +// Source: os/stat_linux.go +func atime(fi os.FileInfo) time.Time { + return timespecToTime(fi.Sys().(*syscall.Stat_t).Atim) +} + +func timespecToTime(ts syscall.Timespec) time.Time { + return time.Unix(ts.Sec, ts.Nsec) +} diff --git a/lib/sshutils/scp/stat_windows.go b/lib/sshutils/scp/stat_windows.go new file mode 100644 index 00000000000..bf8de58c85b --- /dev/null +++ b/lib/sshutils/scp/stat_windows.go @@ -0,0 +1,28 @@ +/* +Copyright 2020 Gravitational, Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package scp + +import ( + "os" + "syscall" + "time" +) + +// Source: os/types_windows.go +func atime(fi os.FileInfo) time.Time { + return time.Unix(0, fi.Sys().(*syscall.Win32FileAttributeData).LastAccessTime.Nanoseconds()) +} diff --git a/lib/web/files.go b/lib/web/files.go index c31d1ebcf80..26badefbdbf 100644 --- a/lib/web/files.go +++ b/lib/web/files.go @@ -17,7 +17,6 @@ limitations under the License. package web import ( - "context" "net/http" "github.com/gravitational/teleport/lib/auth" @@ -104,7 +103,7 @@ func (f *fileTransfer) download(req fileTransferRequest, httpReq *http.Request, return trace.Wrap(err) } - err = tc.ExecuteSCP(context.TODO(), cmd) + err = tc.ExecuteSCP(httpReq.Context(), cmd) if err != nil { return trace.Wrap(err) } @@ -128,7 +127,7 @@ func (f *fileTransfer) upload(req fileTransferRequest, httpReq *http.Request) er return trace.Wrap(err) } - err = tc.ExecuteSCP(context.TODO(), cmd) + err = tc.ExecuteSCP(httpReq.Context(), cmd) if err != nil { return trace.Wrap(err) } diff --git a/tool/teleport/common/teleport.go b/tool/teleport/common/teleport.go index e4778bf1bfc..1aa6c3c85d5 100644 --- a/tool/teleport/common/teleport.go +++ b/tool/teleport/common/teleport.go @@ -145,6 +145,7 @@ func Run(options Options) (executedCommand string, conf *service.Config) { scpc.Flag("v", "verbose mode").Default("false").Short('v').BoolVar(&scpFlags.Verbose) scpc.Flag("r", "recursive mode").Default("false").Short('r').BoolVar(&scpFlags.Recursive) scpc.Flag("d", "directory mode").Short('d').Hidden().BoolVar(&scpFlags.DirectoryMode) + scpc.Flag("preserve", "preserve access and modification times").Short('p').BoolVar(&scpFlags.PreserveAttrs) scpc.Flag("remote-addr", "address of the remote client").StringVar(&scpFlags.RemoteAddr) scpc.Flag("local-addr", "local address which accepted the request").StringVar(&scpFlags.LocalAddr) scpc.Arg("target", "").StringsVar(&scpFlags.Target) diff --git a/tool/tsh/tsh.go b/tool/tsh/tsh.go index f1514d1697a..02bbdb9cb64 100644 --- a/tool/tsh/tsh.go +++ b/tool/tsh/tsh.go @@ -47,6 +47,7 @@ import ( "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/sshutils" + "github.com/gravitational/teleport/lib/sshutils/scp" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/tool/tsh/common" @@ -187,6 +188,9 @@ type CLIConf struct { // terminal. EnableEscapeSequences bool + // PreserveAttrs preserves access/modification times from the original file. + PreserveAttrs bool + // executablePath is the absolute path to the current executable. executablePath string } @@ -291,6 +295,7 @@ func Run(args []string) { scp.Arg("from, to", "Source and destination to copy").Required().StringsVar(&cf.CopySpec) scp.Flag("recursive", "Recursive copy of subdirectories").Short('r').BoolVar(&cf.RecursiveCopy) scp.Flag("port", "Port to connect to on the remote host").Short('P').Int32Var(&cf.NodePort) + scp.Flag("preserve", "Preserves access and modification times from the original file").Short('p').BoolVar(&cf.PreserveAttrs) scp.Flag("quiet", "Quiet mode").Short('q').BoolVar(&cf.Quiet) // ls ls := app.Command("ls", "List remote SSH nodes") @@ -1170,8 +1175,12 @@ func onSCP(cf *CLIConf) { if err != nil { utils.FatalError(err) } + flags := scp.Flags{ + Recursive: cf.RecursiveCopy, + PreserveAttrs: cf.PreserveAttrs, + } err = client.RetryWithRelogin(cf.Context, tc, func() error { - return tc.SCP(context.TODO(), cf.CopySpec, int(cf.NodePort), cf.RecursiveCopy, cf.Quiet) + return tc.SCP(context.TODO(), cf.CopySpec, int(cf.NodePort), flags, cf.Quiet) }) if err != nil { // exit with the same exit status as the failed command: