mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
+10
-14
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user