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
This commit is contained in:
a-palchikov
2021-01-06 13:21:06 +01:00
committed by GitHub
parent a53fc5747e
commit 72630d1df5
12 changed files with 911 additions and 357 deletions
+10 -14
View File
@@ -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 {
+5 -4
View File
@@ -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
+23 -4
View File
@@ -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
}
+25 -6
View File
@@ -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
}
+200 -86
View File
@@ -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<mtime.sec> <mtime.usec> <atime.sec> <atime.usec>
//
// 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<path>.+)`,
// any char including empty which stands for the implicit home directory
`:(?P<path>.*)`,
)
// 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
}
+543 -239
View File
@@ -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
}
+32
View File
@@ -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))
}
+32
View File
@@ -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)
}
+28
View File
@@ -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())
}
+2 -3
View File
@@ -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)
}
+1
View File
@@ -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)
+10 -1
View File
@@ -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: