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
This commit is contained in:
Andrew Lytvynov
2020-05-06 00:03:30 +00:00
committed by Andrew Lytvynov
parent b1eae4ac4c
commit d40013f33b
2 changed files with 86 additions and 35 deletions
+22 -4
View File
@@ -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)
}
+64 -31
View File
@@ -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"))