From d40013f33be0c266d1bdf119dd6452e94de59a2a Mon Sep 17 00:00:00 2001 From: Andrew Lytvynov Date: Tue, 5 May 2020 14:14:08 -0700 Subject: [PATCH] Send SCP error messages to receiver when upload fails This serves two purposes: 1. on upload, the receiver (node) will exit with error and generate an SCP error audit event. 2. on download, the receiver (tsh) will print a meaningful message from the sender (node) Fixes #2861 --- lib/sshutils/scp/scp.go | 26 ++++++++-- lib/sshutils/scp/scp_test.go | 95 ++++++++++++++++++++++++------------ 2 files changed, 86 insertions(+), 35 deletions(-) diff --git a/lib/sshutils/scp/scp.go b/lib/sshutils/scp/scp.go index e60c8dcec23..2c36985757e 100644 --- a/lib/sshutils/scp/scp.go +++ b/lib/sshutils/scp/scp.go @@ -235,11 +235,22 @@ func (cmd *command) GetRemoteShellCmd() (string, error) { return shellCmd, nil } -func (cmd *command) serveSource(ch io.ReadWriter) error { +func (cmd *command) serveSource(ch io.ReadWriter) (retErr error) { + defer func() { + // If anything goes wrong, notify the remote side so it can terminate + // with an error too. + // This is necessary to emit correct audit events (if the remote end is + // emitting them). + if retErr != nil { + cmd.sendErr(ch, retErr) + } + }() + fileInfos := make([]FileInfo, len(cmd.Flags.Target)) 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) } if fileInfo.IsDir() && !cmd.Flags.Recursive { @@ -348,6 +359,13 @@ func (cmd *command) sendFile(r *reader, ch io.ReadWriter, fileInfo FileInfo) err return trace.Wrap(r.read()) } +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) + } +} + // serveSink executes file uploading, when a remote server sends file(s) // via scp func (cmd *command) serveSink(ch io.ReadWriter) error { @@ -408,9 +426,9 @@ func (cmd *command) processCommand(ch io.ReadWriter, st *state, b byte, line str cmd.log.Debugf("[SCP] <- %v %v", string(b), line) switch b { case WarnByte: - return trace.Errorf(line) + return trace.Errorf("error from sender: %q", line) case ErrByte: - return trace.Errorf(line) + return trace.Errorf("error from sender: %q", line) case 'C': f, err := parseNewFile(line) if err != nil { @@ -628,7 +646,7 @@ func (r *reader) read() error { if err := r.s.Err(); err != nil { return trace.Wrap(err) } - return trace.BadParameter(r.s.Text()) + return trace.BadParameter("error from receiver: %q", r.s.Text()) } return trace.BadParameter("unrecognized command: %v", r.b) } diff --git a/lib/sshutils/scp/scp_test.go b/lib/sshutils/scp/scp_test.go index fbeaae6f081..01b7637a6d5 100644 --- a/lib/sshutils/scp/scp_test.go +++ b/lib/sshutils/scp/scp_test.go @@ -107,10 +107,6 @@ func (s *SCPSuite) TestHTTPReceiveFile(c *C) { func (s *SCPSuite) TestSendFile(c *C) { dir := c.MkDir() target := filepath.Join(dir, "target") - contents := []byte("hello, send file!") - - err := ioutil.WriteFile(target, contents, 0666) - c.Assert(err, IsNil) cmd, err := CreateCommand( Config{ @@ -123,8 +119,21 @@ func (s *SCPSuite) TestSendFile(c *C) { ) 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")) @@ -135,10 +144,6 @@ func (s *SCPSuite) TestSendFile(c *C) { func (s *SCPSuite) TestReceiveFile(c *C) { dir := c.MkDir() source := filepath.Join(dir, "target") - contents := []byte("hello, file contents!") - - err := ioutil.WriteFile(source, contents, 0666) - c.Assert(err, IsNil) outDir := c.MkDir() + "/" cmd, err := CreateCommand(Config{ @@ -150,6 +155,19 @@ func (s *SCPSuite) TestReceiveFile(c *C) { }) c.Assert(err, IsNil) + // Source file is missing, expect an error. + err = runSCP(cmd, "scp", "-v", "-f", source) + 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, file contents!") + err = ioutil.WriteFile(source, contents, 0666) + c.Assert(err, IsNil) + + // Source file is present, should send fine. err = runSCP(cmd, "scp", "-v", "-f", source) c.Assert(err, IsNil) @@ -159,16 +177,7 @@ func (s *SCPSuite) TestReceiveFile(c *C) { } func (s *SCPSuite) TestSendDir(c *C) { - dir := c.MkDir() - c.Assert(os.Mkdir(filepath.Join(dir, "target_dir"), 0777), IsNil) - - err := ioutil.WriteFile( - filepath.Join(dir, "target_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) + dir := filepath.Join(c.MkDir(), "target_dir") cmd, err := CreateCommand(Config{ User: "test-user", @@ -180,12 +189,30 @@ func (s *SCPSuite) TestSendDir(c *C) { }) 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, "target_dir", "target1")) + bytes, err := ioutil.ReadFile(filepath.Join(outDir, name, "nested_dir", "target1")) c.Assert(err, IsNil) c.Assert(string(bytes), Equals, string("file 1")) @@ -195,16 +222,7 @@ func (s *SCPSuite) TestSendDir(c *C) { } func (s *SCPSuite) TestReceiveDir(c *C) { - dir := c.MkDir() - c.Assert(os.Mkdir(filepath.Join(dir, "target_dir"), 0777), IsNil) - - err := ioutil.WriteFile( - filepath.Join(dir, "target_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) + dir := filepath.Join(c.MkDir(), "target_dir") outDir := c.MkDir() + "/" cmd, err := CreateCommand(Config{ @@ -217,12 +235,27 @@ func (s *SCPSuite) TestReceiveDir(c *C) { }) c.Assert(err, IsNil) + // Source directory is missing, expect an error. + err = runSCP(cmd, "scp", "-v", "-r", "-f", dir) + c.Assert(err, NotNil) + + // 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", "-f", dir) c.Assert(err, IsNil) - time.Sleep(time.Millisecond * 300) name := filepath.Base(dir) - bytes, err := ioutil.ReadFile(filepath.Join(outDir, name, "target_dir", "target1")) + bytes, err := ioutil.ReadFile(filepath.Join(outDir, name, "nested_dir", "target1")) c.Assert(err, IsNil) c.Assert(string(bytes), Equals, string("file 1"))