From bcffdcbfd68892a047edaba9d49af034c85974ee Mon Sep 17 00:00:00 2001 From: Edoardo Spadolini Date: Wed, 15 Apr 2026 15:28:27 +0200 Subject: [PATCH] session-helper: move the sftp subcommand to reexec.RunAndExit (#65392) * Split the approver map away from FileTransferRequest * Move or copy files in session/reexec/sftputils and session/reexec/reexecsftp * Split up copied files and clean up the new session packages * Use custom types for SFTP audit log events * Clean up imports and run sftp in RunAndExit * Finish renaming types and functions * Streamline reexec in main and tests * Add test with the legacy SFTP event implementation --- constants.go | 3 - integration/helpers/testmain.go | 10 +- lib/srv/ctx.go | 7 +- lib/srv/exec_test.go | 7 +- lib/srv/forward/sftp.go | 29 +- lib/srv/forward/sftp_test.go | 2 +- lib/srv/regular/sftp.go | 54 ++- lib/srv/regular/sshserver_test.go | 5 +- lib/srv/sess.go | 38 +- lib/srv/sess_test.go | 31 +- lib/sshutils/sftp/events.go | 152 +++++++ lib/sshutils/sftp/events_test.go | 229 ++++++++++ lib/sshutils/sftp/http.go | 7 +- lib/sshutils/sftp/remote.go | 7 +- lib/sshutils/sftp/sftp.go | 334 +------------- lib/sshutils/sftp/sftp_test.go | 405 +---------------- lib/sshutils/sftp/utils.go | 199 --------- lib/tbot/tbot_test.go | 2 + .../clusters/cluster_file_transfer.go | 3 +- lib/web/apiserver_test.go | 7 +- lib/web/files.go | 3 +- session/reexec/exec_test.go | 5 +- session/reexec/reexec.go | 40 +- session/reexec/reexecsftp/sess.go | 32 ++ .../reexec/reexecsftp}/sftp.go | 70 +-- .../sftp => session/sftputils}/local.go | 56 ++- .../sftp => session/sftputils}/local_test.go | 2 +- session/sftputils/sftp.go | 349 +++++++++++++++ session/sftputils/sftp_test.go | 419 ++++++++++++++++++ session/sftputils/utils.go | 233 ++++++++++ tool/teleport/common/teleport.go | 10 +- tool/teleport/main.go | 3 + tool/teleport/testenv/test_server.go | 10 - tool/tsh/common/tsh_test.go | 6 +- 34 files changed, 1625 insertions(+), 1144 deletions(-) create mode 100644 lib/sshutils/sftp/events.go create mode 100644 lib/sshutils/sftp/events_test.go create mode 100644 session/reexec/reexecsftp/sess.go rename {tool/teleport/common => session/reexec/reexecsftp}/sftp.go (82%) rename {lib/sshutils/sftp => session/sftputils}/local.go (61%) rename {lib/sshutils/sftp => session/sftputils}/local_test.go (99%) create mode 100644 session/sftputils/sftp.go create mode 100644 session/sftputils/sftp_test.go create mode 100644 session/sftputils/utils.go diff --git a/constants.go b/constants.go index d1e36ab2d87..8c89395e136 100644 --- a/constants.go +++ b/constants.go @@ -171,9 +171,6 @@ const ( // ComponentSubsystemProxy is the proxy subsystem. ComponentSubsystemProxy = "subsystem:proxy" - // ComponentSubsystemSFTP is the SFTP subsystem. - ComponentSubsystemSFTP = "subsystem:sftp" - // ComponentLocalTerm is a terminal on a regular SSH node. ComponentLocalTerm = "term:local" diff --git a/integration/helpers/testmain.go b/integration/helpers/testmain.go index 07094cf7b1b..e93506b92a3 100644 --- a/integration/helpers/testmain.go +++ b/integration/helpers/testmain.go @@ -28,25 +28,19 @@ import ( "github.com/gravitational/teleport/lib/modules" "github.com/gravitational/teleport/lib/utils/log/logtest" "github.com/gravitational/teleport/session/reexec" - "github.com/gravitational/teleport/tool/teleport/common" ) // TestMainImplementation will re-execute Teleport to run a command if "exec" is passed to // it as an argument. Otherwise, it will run tests as normal. func TestMainImplementation(m *testing.M) { + reexec.MaybeReexec() + logtest.InitLogger(testing.Verbose) ctx, cancel := context.WithCancel(context.Background()) cryptosuitestest.PrecomputeRSAKeys(ctx) SetTestTimeouts(3 * time.Second) modules.SetInsecureTestMode(true) - // If the test is re-executing itself, execute the command that comes over - // the pipe. - if reexec.IsReexec() { - defer cancel() - common.Run(common.Options{Args: os.Args[1:]}) - return - } // Otherwise run tests as normal. exitCode := m.Run() diff --git a/lib/srv/ctx.go b/lib/srv/ctx.go index cf8499b4cb0..02fe8e1b865 100644 --- a/lib/srv/ctx.go +++ b/lib/srv/ctx.go @@ -62,6 +62,7 @@ import ( "github.com/gravitational/teleport/session/networking/x11" "github.com/gravitational/teleport/session/pam/pamcfg" "github.com/gravitational/teleport/session/reexec" + "github.com/gravitational/teleport/session/reexec/reexecsftp" ) var ctxID int32 @@ -461,7 +462,7 @@ type ServerContext struct { // approvedFileReq is an approved file transfer request that will only be // set when the session's pending file transfer request is approved. - approvedFileReq *FileTransferRequest + approvedFileReq *reexecsftp.FileTransferRequest } // NewServerContext creates a new *ServerContext which is used to pass and @@ -1374,7 +1375,7 @@ func (c *ServerContext) GetPortForwardEvent(evType, code, addr string) apievents } } -func (c *ServerContext) setApprovedFileTransferRequest(req *FileTransferRequest) { +func (c *ServerContext) setApprovedFileTransferRequest(req *reexecsftp.FileTransferRequest) { c.mu.Lock() c.approvedFileReq = req c.mu.Unlock() @@ -1384,7 +1385,7 @@ func (c *ServerContext) setApprovedFileTransferRequest(req *FileTransferRequest) // request for this session if there is one present. Note that if an // approved request is returned future calls to this method will return // nil to prevent an approved request getting reused incorrectly. -func (c *ServerContext) ConsumeApprovedFileTransferRequest() *FileTransferRequest { +func (c *ServerContext) ConsumeApprovedFileTransferRequest() *reexecsftp.FileTransferRequest { c.mu.Lock() defer c.mu.Unlock() diff --git a/lib/srv/exec_test.go b/lib/srv/exec_test.go index 6ae72308564..52a8170a073 100644 --- a/lib/srv/exec_test.go +++ b/lib/srv/exec_test.go @@ -40,14 +40,9 @@ import ( // TestMain will re-execute Teleport to run a command if "exec" is passed to // it as an argument. Otherwise, it will run tests as normal. func TestMain(m *testing.M) { + reexec.MaybeReexec() logtest.InitLogger(testing.Verbose) modules.SetInsecureTestMode(true) - // If the test is re-executing itself, execute the command that comes over - // the pipe. - if reexec.IsReexec() { - reexec.RunAndExit(os.Args[1]) - return - } // Otherwise run tests as normal. code := m.Run() diff --git a/lib/srv/forward/sftp.go b/lib/srv/forward/sftp.go index 31ce871b85b..c54e065409c 100644 --- a/lib/srv/forward/sftp.go +++ b/lib/srv/forward/sftp.go @@ -33,6 +33,7 @@ import ( "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/srv" sftputils "github.com/gravitational/teleport/lib/sshutils/sftp" + sessionsftputils "github.com/gravitational/teleport/session/sftputils" ) // SFTPProxy proxies an SFTP session and emits audit events for the handled @@ -127,11 +128,11 @@ func (p *SFTPProxy) Close() error { type proxyHandlers struct { auditContext sftpAuditContext - remoteFS sftputils.FileSystem + remoteFS sessionsftputils.FileSystem logger *slog.Logger fileMtx sync.Mutex - files []*sftputils.TrackedFile + files []*sessionsftputils.TrackedFile } type sftpAuditContext interface { @@ -190,15 +191,15 @@ func (h *proxyHandlers) OpenFile(req *sftp.Request) (_ sftp.WriterAtReaderAt, re return nil, os.ErrInvalid } - f, err := h.remoteFS.OpenFile(req.Filepath, sftputils.ParseFlags(req)) + f, err := h.remoteFS.OpenFile(req.Filepath, sessionsftputils.ParseFlags(req)) if err != nil { return nil, err } return h.trackFile(f), nil } -func (h *proxyHandlers) trackFile(f sftputils.File) sftp.WriterAtReaderAt { - trackFile := &sftputils.TrackedFile{File: f} +func (h *proxyHandlers) trackFile(f sessionsftputils.File) sftp.WriterAtReaderAt { + trackFile := &sessionsftputils.TrackedFile{File: f} h.fileMtx.Lock() defer h.fileMtx.Unlock() h.files = append(h.files, trackFile) @@ -212,17 +213,17 @@ func (h *proxyHandlers) Filecmd(req *sftp.Request) (err error) { h.sendSFTPEvent(req, err) } }() - return sftputils.HandleFilecmd(req, h.remoteFS) + return sessionsftputils.HandleFilecmd(req, h.remoteFS) } // Filelist handles listing info about files. func (h *proxyHandlers) Filelist(req *sftp.Request) (_ sftp.ListerAt, err error) { defer func() { - if req.Method == sftputils.MethodList { + if req.Method == sessionsftputils.MethodList { h.sendSFTPEvent(req, err) } }() - lister, err := sftputils.HandleFilelist(req, h.remoteFS) + lister, err := sessionsftputils.HandleFilelist(req, h.remoteFS) if err != nil { return nil, err } @@ -241,11 +242,17 @@ func (h *proxyHandlers) sendSFTPEvent(req *sftp.Request, reqErr error) { h.logger.WarnContext(req.Context(), "Unable to get working directory", "error", err) // Emit event without working directory. } - event, err := sftputils.ParseSFTPEvent(req, wd, reqErr) + sftpEvent, err := sessionsftputils.ParseSFTPEvent(req, wd, reqErr) if err != nil { - h.logger.WarnContext(req.Context(), "Unknown SFTP request", "request", req.Method) + h.logger.WarnContext(req.Context(), "Failed to convert SFTP event into an audit log event", "request", req.Method, "error", err) return - } else if reqErr != nil { + } + event, err := sftputils.SFTPEventToProto(sftpEvent) + if err != nil { + h.logger.WarnContext(req.Context(), "Failed to convert SFTP event into an audit log event", "request", req.Method, "error", err) + return + } + if reqErr != nil { h.logger.DebugContext(req.Context(), "failed handling SFTP request", "request", req.Method, "error", reqErr) } event.ServerMetadata = h.auditContext.ServerMetadata() diff --git a/lib/srv/forward/sftp_test.go b/lib/srv/forward/sftp_test.go index 2caef2406e5..f671f462792 100644 --- a/lib/srv/forward/sftp_test.go +++ b/lib/srv/forward/sftp_test.go @@ -29,7 +29,7 @@ import ( "github.com/stretchr/testify/require" apievents "github.com/gravitational/teleport/api/types/events" - sftputils "github.com/gravitational/teleport/lib/sshutils/sftp" + "github.com/gravitational/teleport/session/sftputils" ) func TestSFTPProxyServeClosesRemoteFilesystem(t *testing.T) { diff --git a/lib/srv/regular/sftp.go b/lib/srv/regular/sftp.go index 94a2606a856..9abc8e4847f 100644 --- a/lib/srv/regular/sftp.go +++ b/lib/srv/regular/sftp.go @@ -19,7 +19,6 @@ package regular import ( - "bufio" "context" "encoding/json" "errors" @@ -27,11 +26,9 @@ import ( "log/slog" "os" "os/exec" - "strings" "sync" "time" - "github.com/gogo/protobuf/jsonpb" //nolint:depguard // needed for backwards compatibility "github.com/gravitational/trace" "golang.org/x/crypto/ssh" @@ -40,14 +37,17 @@ import ( "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/srv" "github.com/gravitational/teleport/lib/sshutils/reexec" + sftputils "github.com/gravitational/teleport/lib/sshutils/sftp" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/session/reexec/reexecconstants" + "github.com/gravitational/teleport/session/reexec/reexecsftp" + sessionsftputils "github.com/gravitational/teleport/session/sftputils" ) type sftpSubsys struct { logger *slog.Logger - fileTransferReq *srv.FileTransferRequest + fileTransferReq *reexecsftp.FileTransferRequest sftpCmd *exec.Cmd serverCtx *srv.ServerContext @@ -57,9 +57,9 @@ type sftpSubsys struct { waitForOutputStreams sync.WaitGroup } -func newSFTPSubsys(fileTransferReq *srv.FileTransferRequest) (*sftpSubsys, error) { +func newSFTPSubsys(fileTransferReq *reexecsftp.FileTransferRequest) (*sftpSubsys, error) { return &sftpSubsys{ - logger: slog.With(teleport.ComponentKey, teleport.ComponentSubsystemSFTP), + logger: slog.With(teleport.ComponentKey, "subsystem:sftp"), fileTransferReq: fileTransferReq, }, nil } @@ -207,44 +207,40 @@ func (s *sftpSubsys) Start(ctx context.Context, LocalAddr: serverConn.LocalAddr().String(), } - r := bufio.NewReader(auditPipeOut) + dec := json.NewDecoder(auditPipeOut) for { - // Read up to a NULL byte, the child process uses this to - // delimit audit events - eventStr, err := r.ReadString(0x0) - if err != nil { + var ev sessionsftputils.Event + if err := dec.Decode(&ev); err != nil { if !errors.Is(err, io.EOF) { s.logger.WarnContext(ctx, "Failed to read SFTP event", "error", err) } return } - var oneOfEvent apievents.OneOf - err = (&jsonpb.Unmarshaler{}).Unmarshal(strings.NewReader(eventStr[:len(eventStr)-1]), &oneOfEvent) - if err != nil { - s.logger.WarnContext(ctx, "Failed to unmarshal SFTP event", "error", err) - continue - } - event, err := apievents.FromOneOf(oneOfEvent) - if err != nil { - s.logger.WarnContext(ctx, "Failed to convert SFTP event from OneOf", "error", err) - continue - } - - event.SetClusterName(serverCtx.ClusterName) - switch e := event.(type) { - case *apievents.SFTP: + var event apievents.AuditEvent + if ev.SFTP != nil { + e, err := sftputils.SFTPEventToProto(ev.SFTP) + if err != nil { + s.logger.WarnContext(ctx, "Failed to convert SFTP event", "error", err) + continue + } + e.SetClusterName(serverCtx.ClusterName) e.ServerMetadata = serverMeta e.SessionMetadata = sessionMeta e.UserMetadata = userMeta e.ConnectionMetadata = connectionMeta - case *apievents.SFTPSummary: + event = e + } else if ev.Summary != nil { + e := sftputils.SFTPSummaryEventToProto(ev.Summary) + e.SetClusterName(serverCtx.ClusterName) e.ServerMetadata = serverMeta e.SessionMetadata = sessionMeta e.UserMetadata = userMeta e.ConnectionMetadata = connectionMeta - default: - s.logger.WarnContext(ctx, "Unknown event type received from SFTP server process", "error", err, "event_type", event.GetType()) + event = e + } else { + s.logger.WarnContext(ctx, "Unknown event type received from SFTP server process") + continue } if err := serverCtx.GetServer().EmitAuditEvent(ctx, event); err != nil { diff --git a/lib/srv/regular/sshserver_test.go b/lib/srv/regular/sshserver_test.go index 7b90bd82170..cfa83c41415 100644 --- a/lib/srv/regular/sshserver_test.go +++ b/lib/srv/regular/sshserver_test.go @@ -100,12 +100,9 @@ var wildcardAllow = types.Labels{ // TestMain will re-execute Teleport to run a command if "exec" is passed to // it as an argument. Otherwise it will run tests as normal. func TestMain(m *testing.M) { + reexec.MaybeReexec() logtest.InitLogger(testing.Verbose) modules.SetInsecureTestMode(true) - if reexec.IsReexec() { - reexec.RunAndExit(os.Args[1]) - return - } code := m.Run() os.Exit(code) diff --git a/lib/srv/sess.go b/lib/srv/sess.go index a9078c19f91..abc1f038cbf 100644 --- a/lib/srv/sess.go +++ b/lib/srv/sess.go @@ -56,6 +56,7 @@ import ( "github.com/gravitational/teleport/lib/services" rsession "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/utils" + "github.com/gravitational/teleport/session/reexec/reexecsftp" ) const sessionRecorderID = "session-recorder" @@ -520,7 +521,7 @@ func (s *SessionRegistry) isApprovedFileTransfer(scx *ServerContext) (bool, erro return false, trace.Wrap(err) } if approved { - scx.setApprovedFileTransferRequest(sess.fileTransferReq) + scx.setApprovedFileTransferRequest(&sess.fileTransferReq.FileTransferRequest) sess.fileTransferReq = nil } @@ -545,7 +546,7 @@ const ( // notifyFileTransferRequestUnderLock is called to notify all members of a party that a file transfer request has been created/approved/denied. // The notification is a global ssh request and requires the client to update its UI state accordingly. -func (s *SessionRegistry) notifyFileTransferRequestUnderLock(req *FileTransferRequest, res FileTransferRequestEvent, scx *ServerContext) error { +func (s *SessionRegistry) notifyFileTransferRequestUnderLock(req *fileTransferRequestWithApprovers, res FileTransferRequestEvent, scx *ServerContext) error { session := scx.getSession() if session == nil { s.logger.DebugContext( @@ -744,7 +745,7 @@ type session struct { // fileTransferReq a pending file transfer request for this session. // If the request is denied or approved it should be set to nil to // prevent its reuse. - fileTransferReq *FileTransferRequest + fileTransferReq *fileTransferRequestWithApprovers io *TermManager inWriter io.WriteCloser @@ -1840,24 +1841,13 @@ func (s *session) checkPresence(ctx context.Context) error { return nil } -// FileTransferRequest is a request to upload or download a file from a node. -type FileTransferRequest struct { - // ID is a UUID that uniquely identifies a file transfer request - // and is unlikely to collide with another file transfer request - ID string - // Requester is the Teleport User that requested the file transfer - Requester string - // Download is true if the request is a download, false if its an upload - Download bool - // Filename is the name of the file to upload. - Filename string - // Location of the requested download or where a file will be uploaded - Location string +type fileTransferRequestWithApprovers struct { + reexecsftp.FileTransferRequest // approvers is a list of participants of moderator or peer type that have approved the request approvers map[string]*party } -func (s *session) checkIfFileTransferApproved(req *FileTransferRequest) (bool, error) { +func (s *session) checkIfFileTransferApproved(req *fileTransferRequestWithApprovers) (bool, error) { var participants []moderation.SessionAccessContext for _, party := range req.approvers { @@ -1898,12 +1888,14 @@ func (s *session) addFileTransferRequest(params *rsession.FileTransferRequestPar return trace.BadParameter("no source file is set for the upload") } - s.fileTransferReq = &FileTransferRequest{ - ID: uuid.New().String(), - Requester: params.Requester, - Location: params.Location, - Filename: params.Filename, - Download: params.Download, + s.fileTransferReq = &fileTransferRequestWithApprovers{ + FileTransferRequest: reexecsftp.FileTransferRequest{ + ID: uuid.New().String(), + Requester: params.Requester, + Location: params.Location, + Filename: params.Filename, + Download: params.Download, + }, approvers: make(map[string]*party), } diff --git a/lib/srv/sess_test.go b/lib/srv/sess_test.go index f785653ae42..8238c024cdf 100644 --- a/lib/srv/sess_test.go +++ b/lib/srv/sess_test.go @@ -53,6 +53,7 @@ import ( rsession "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/sshutils/sftp" "github.com/gravitational/teleport/lib/utils/log/logtest" + "github.com/gravitational/teleport/session/reexec/reexecsftp" ) func TestIsApprovedFileTransfer(t *testing.T) { @@ -106,7 +107,7 @@ func TestIsApprovedFileTransfer(t *testing.T) { name string expectedResult bool expectedError string - req *FileTransferRequest + req *fileTransferRequestWithApprovers reqID string location string }{ @@ -122,9 +123,11 @@ func TestIsApprovedFileTransfer(t *testing.T) { expectedResult: false, expectedError: "Teleport user does not match original requester", reqID: "123", - req: &FileTransferRequest{ - ID: "123", - Requester: "michael", + req: &fileTransferRequestWithApprovers{ + FileTransferRequest: reexecsftp.FileTransferRequest{ + ID: "123", + Requester: "michael", + }, approvers: make(map[string]*party), }, }, @@ -134,11 +137,13 @@ func TestIsApprovedFileTransfer(t *testing.T) { expectedError: "requested destination path does not match the current request", reqID: "123", location: "~/Downloads", - req: &FileTransferRequest{ - ID: "123", - Requester: "teleportUser", + req: &fileTransferRequestWithApprovers{ + FileTransferRequest: reexecsftp.FileTransferRequest{ + ID: "123", + Requester: "teleportUser", + Location: "~/badlocation", + }, approvers: make(map[string]*party), - Location: "~/badlocation", }, }, { @@ -147,11 +152,13 @@ func TestIsApprovedFileTransfer(t *testing.T) { expectedError: "", reqID: "123", location: "~/Downloads", - req: &FileTransferRequest{ - ID: "123", - Requester: "teleportUser", + req: &fileTransferRequestWithApprovers{ + FileTransferRequest: reexecsftp.FileTransferRequest{ + ID: "123", + Requester: "teleportUser", + Location: "~/Downloads", + }, approvers: approvers, - Location: "~/Downloads", }, }, } diff --git a/lib/sshutils/sftp/events.go b/lib/sshutils/sftp/events.go new file mode 100644 index 00000000000..91c4552e340 --- /dev/null +++ b/lib/sshutils/sftp/events.go @@ -0,0 +1,152 @@ +// Teleport +// Copyright (C) 2026 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package sftp + +import ( + "time" + + "github.com/gravitational/trace" + + apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/lib/events" + "github.com/gravitational/teleport/session/sftputils" +) + +// SFTPSummaryEventToProto converts a sftp event from the format generated by +// the SFTP server process to its matching audit log struct. +func SFTPEventToProto(ev *sftputils.SFTPEvent) (*apievents.SFTP, error) { + event := &apievents.SFTP{ + Metadata: apievents.Metadata{ + Type: events.SFTPEvent, + Time: time.Unix(0, ev.Time), + }, + } + + switch ev.Method { + case sftputils.MethodOpen, sftputils.MethodGet, sftputils.MethodPut: + if ev.Error == "" { + event.Code = events.SFTPOpenCode + } else { + event.Code = events.SFTPOpenFailureCode + } + event.Action = apievents.SFTPAction_OPEN + case sftputils.MethodSetStat: + if ev.Error == "" { + event.Code = events.SFTPSetstatCode + } else { + event.Code = events.SFTPSetstatFailureCode + } + event.Action = apievents.SFTPAction_SETSTAT + case sftputils.MethodList: + if ev.Error == "" { + event.Code = events.SFTPReaddirCode + } else { + event.Code = events.SFTPReaddirFailureCode + } + event.Action = apievents.SFTPAction_READDIR + case sftputils.MethodRemove: + if ev.Error == "" { + event.Code = events.SFTPRemoveCode + } else { + event.Code = events.SFTPRemoveFailureCode + } + event.Action = apievents.SFTPAction_REMOVE + case sftputils.MethodMkdir: + if ev.Error == "" { + event.Code = events.SFTPMkdirCode + } else { + event.Code = events.SFTPMkdirFailureCode + } + event.Action = apievents.SFTPAction_MKDIR + case sftputils.MethodRmdir: + if ev.Error == "" { + event.Code = events.SFTPRmdirCode + } else { + event.Code = events.SFTPRmdirFailureCode + } + event.Action = apievents.SFTPAction_RMDIR + case sftputils.MethodRename: + if ev.Error == "" { + event.Code = events.SFTPRenameCode + } else { + event.Code = events.SFTPRenameFailureCode + } + event.Action = apievents.SFTPAction_RENAME + case sftputils.MethodSymlink: + if ev.Error == "" { + event.Code = events.SFTPSymlinkCode + } else { + event.Code = events.SFTPSymlinkFailureCode + } + event.Action = apievents.SFTPAction_SYMLINK + case sftputils.MethodLink: + if ev.Error == "" { + event.Code = events.SFTPLinkCode + } else { + event.Code = events.SFTPLinkFailureCode + } + event.Action = apievents.SFTPAction_LINK + default: + return nil, trace.BadParameter("unknown SFTP request %+q", ev.Method) + } + + event.Path = ev.Path + event.TargetPath = ev.Target + event.Flags = ev.Flags + event.WorkingDirectory = ev.WorkDir + + if ev.Attrs != nil { + event.Attributes = new(apievents.SFTPAttributes) + if ev.Attrs.Atime != nil { + t := time.Unix(int64(*ev.Attrs.Atime), 0) + event.Attributes.AccessTime = &t + } + if ev.Attrs.Mtime != nil { + t := time.Unix(int64(*ev.Attrs.Mtime), 0) + event.Attributes.ModificationTime = &t + } + event.Attributes.Permissions = ev.Attrs.Perms + event.Attributes.FileSize = ev.Attrs.Size + event.Attributes.UID = ev.Attrs.UID + event.Attributes.GID = ev.Attrs.GID + } + + event.Error = ev.Error + + return event, nil +} + +// SFTPSummaryEventToProto converts a sftp_summary event from the format +// generated by the SFTP server process to its matching audit log struct. +func SFTPSummaryEventToProto(ev *sftputils.SFTPSummaryEvent) *apievents.SFTPSummary { + event := &apievents.SFTPSummary{ + Metadata: apievents.Metadata{ + Type: events.SFTPSummaryEvent, + Code: events.SFTPSummaryCode, + Time: time.Now(), + }, + FileTransferStats: make([]*apievents.FileTransferStat, 0, len(ev.Stats)), + } + for _, stat := range ev.Stats { + event.FileTransferStats = append(event.FileTransferStats, &apievents.FileTransferStat{ + Path: stat.Path, + BytesRead: stat.Read, + BytesWritten: stat.Written, + }) + } + return event +} diff --git a/lib/sshutils/sftp/events_test.go b/lib/sshutils/sftp/events_test.go new file mode 100644 index 00000000000..7bb34ccf4a1 --- /dev/null +++ b/lib/sshutils/sftp/events_test.go @@ -0,0 +1,229 @@ +// Teleport +// Copyright (C) 2026 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package sftp + +import ( + "errors" + "io/fs" + "os" + "testing" + "time" + + "github.com/google/go-cmp/cmp" + "github.com/gravitational/trace" + "github.com/pkg/sftp" + "github.com/stretchr/testify/require" + "google.golang.org/protobuf/testing/protocmp" + + apievents "github.com/gravitational/teleport/api/types/events" + "github.com/gravitational/teleport/lib/events" + "github.com/gravitational/teleport/session/sftputils" +) + +// legacyParseSFTPEvent is the original implementation of event conversion from +// [sftp.Request] to audit log event, before the change to custom types used by +// the SFTP server process. +func legacyParseSFTPEvent(req *sftp.Request, workingDirectory string, reqErr error) (*apievents.SFTP, error) { + event := &apievents.SFTP{ + Metadata: apievents.Metadata{ + Type: events.SFTPEvent, + Time: time.Now(), + }, + } + + switch req.Method { + case sftputils.MethodOpen, sftputils.MethodGet, sftputils.MethodPut: + if reqErr == nil { + event.Code = events.SFTPOpenCode + } else { + event.Code = events.SFTPOpenFailureCode + } + event.Action = apievents.SFTPAction_OPEN + case sftputils.MethodSetStat: + if reqErr == nil { + event.Code = events.SFTPSetstatCode + } else { + event.Code = events.SFTPSetstatFailureCode + } + event.Action = apievents.SFTPAction_SETSTAT + case sftputils.MethodList: + if reqErr == nil { + event.Code = events.SFTPReaddirCode + } else { + event.Code = events.SFTPReaddirFailureCode + } + event.Action = apievents.SFTPAction_READDIR + case sftputils.MethodRemove: + if reqErr == nil { + event.Code = events.SFTPRemoveCode + } else { + event.Code = events.SFTPRemoveFailureCode + } + event.Action = apievents.SFTPAction_REMOVE + case sftputils.MethodMkdir: + if reqErr == nil { + event.Code = events.SFTPMkdirCode + } else { + event.Code = events.SFTPMkdirFailureCode + } + event.Action = apievents.SFTPAction_MKDIR + case sftputils.MethodRmdir: + if reqErr == nil { + event.Code = events.SFTPRmdirCode + } else { + event.Code = events.SFTPRmdirFailureCode + } + event.Action = apievents.SFTPAction_RMDIR + case sftputils.MethodRename: + if reqErr == nil { + event.Code = events.SFTPRenameCode + } else { + event.Code = events.SFTPRenameFailureCode + } + event.Action = apievents.SFTPAction_RENAME + case sftputils.MethodSymlink: + if reqErr == nil { + event.Code = events.SFTPSymlinkCode + } else { + event.Code = events.SFTPSymlinkFailureCode + } + event.Action = apievents.SFTPAction_SYMLINK + case sftputils.MethodLink: + if reqErr == nil { + event.Code = events.SFTPLinkCode + } else { + event.Code = events.SFTPLinkFailureCode + } + event.Action = apievents.SFTPAction_LINK + default: + return nil, trace.BadParameter("unknown SFTP request %q", req.Method) + } + + event.Path = req.Filepath + event.TargetPath = req.Target + event.Flags = req.Flags + event.WorkingDirectory = workingDirectory + if req.Method == sftputils.MethodSetStat { + attrFlags := req.AttrFlags() + attrs := req.Attributes() + event.Attributes = new(apievents.SFTPAttributes) + + if attrFlags.Acmodtime { + atime := time.Unix(int64(attrs.Atime), 0) + mtime := time.Unix(int64(attrs.Mtime), 0) + event.Attributes.AccessTime = &atime + event.Attributes.ModificationTime = &mtime + } + if attrFlags.Permissions { + perms := uint32(attrs.FileMode().Perm()) + event.Attributes.Permissions = &perms + } + if attrFlags.Size { + event.Attributes.FileSize = &attrs.Size + } + if attrFlags.UidGid { + event.Attributes.UID = &attrs.UID + event.Attributes.GID = &attrs.GID + } + } + if reqErr != nil { + // If possible, strip the filename from the error message. The + // path will be included in audit events already, no need to + // make the error message longer than it needs to be. + var pathErr *fs.PathError + var linkErr *os.LinkError + if errors.As(reqErr, &pathErr) { + event.Error = pathErr.Err.Error() + } else if errors.As(reqErr, &linkErr) { + event.Error = linkErr.Err.Error() + } else { + event.Error = reqErr.Error() + } + } + + return event, nil +} + +func TestSFTPEventMatchesLegacy(t *testing.T) { + // sftp protocol constants + const ( + sshFileXferAttrSize = 0x00000001 + sshFileXferAttrUIDGID = 0x00000002 + sshFileXferAttrPermissions = 0x00000004 + sshFileXferAttrACmodTime = 0x00000008 + ) + + inputs := []struct { + req *sftp.Request + workingDirectory string + reqErr error + }{ + {&sftp.Request{ + Method: sftputils.MethodGet, + Filepath: "/fp", + }, "/mywd", nil}, + {&sftp.Request{ + Method: sftputils.MethodPut, + Filepath: "/fp", + Flags: 42, + }, "/mywd", nil}, + {&sftp.Request{ + Method: sftputils.MethodPut, + Filepath: "/fp", + }, "/mywd", &fs.PathError{Path: "/fp", Err: errors.New("lmao")}}, + {&sftp.Request{ + Method: sftputils.MethodRemove, + Filepath: "/fp", + }, "/mywd", nil}, + {&sftp.Request{ + Method: sftputils.MethodLink, + Filepath: "/fp", + Target: "/fp2", + }, "/mywd", &os.LinkError{Old: "/fp", New: "/fp2", Err: errors.New("lmao")}}, + {&sftp.Request{ + Method: sftputils.MethodSetStat, + Filepath: "/fp", + Flags: sshFileXferAttrACmodTime, + Attrs: []byte{0x1, 0x23, 0x45, 0x67, 0x12, 0x34, 0x56, 0x78}, + }, "/mywd", nil}, + {&sftp.Request{ + Method: sftputils.MethodSetStat, + Filepath: "/fp", + Flags: sshFileXferAttrSize, + Attrs: []byte{0x1, 0x23, 0x45, 0x67, 0x12, 0x34, 0x56, 0x78}, + }, "/mywd", nil}, + {&sftp.Request{ + Method: sftputils.MethodSetStat, + Filepath: "/fp", + Flags: sshFileXferAttrPermissions, + Attrs: []byte{0, 0, 0o7, 0o55}, + }, "/mywd", nil}, + } + + for _, input := range inputs { + legacyEvent, err := legacyParseSFTPEvent(input.req, input.workingDirectory, input.reqErr) + require.NoError(t, err) + + sftpEvent, err := sftputils.ParseSFTPEvent(input.req, input.workingDirectory, input.reqErr) + require.NoError(t, err) + + newEvent, err := SFTPEventToProto(sftpEvent) + require.NoError(t, err) + + require.Empty(t, cmp.Diff(legacyEvent, newEvent, protocmp.Transform())) + } +} diff --git a/lib/sshutils/sftp/http.go b/lib/sshutils/sftp/http.go index 806afb2a671..2564e122eb4 100644 --- a/lib/sshutils/sftp/http.go +++ b/lib/sshutils/sftp/http.go @@ -32,6 +32,7 @@ import ( "github.com/gravitational/trace" "github.com/gravitational/teleport/lib/httplib" + "github.com/gravitational/teleport/session/sftputils" ) const ( @@ -71,7 +72,7 @@ func (h *httpFS) ReadDir(_ string) ([]fs.FileInfo, error) { return nil, errDirsNotSupported } -func (h *httpFS) Open(path string) (File, error) { +func (h *httpFS) Open(path string) (sftputils.File, error) { if h.reader == nil { return nil, trace.BadParameter("missing reader") } @@ -85,7 +86,7 @@ func (h *httpFS) Open(path string) (File, error) { }, nil } -func (h *httpFS) Create(p string, size int64) (File, error) { +func (h *httpFS) Create(p string, size int64) (sftputils.File, error) { filename := path.Base(p) contentLength := strconv.FormatInt(size, 10) header := h.writer.Header() @@ -106,7 +107,7 @@ func (h *httpFS) Create(p string, size int64) (File, error) { }, nil } -func (h *httpFS) OpenFile(p string, flags int) (File, error) { +func (h *httpFS) OpenFile(p string, flags int) (sftputils.File, error) { switch flags & 3 { case os.O_RDWR: return nil, trace.BadParameter("read-write files not supported for http") diff --git a/lib/sshutils/sftp/remote.go b/lib/sshutils/sftp/remote.go index 81c7257ddd1..68253ca68ee 100644 --- a/lib/sshutils/sftp/remote.go +++ b/lib/sshutils/sftp/remote.go @@ -30,6 +30,7 @@ import ( "github.com/gravitational/teleport" tracessh "github.com/gravitational/teleport/api/observability/tracing/ssh" + "github.com/gravitational/teleport/session/sftputils" ) // RemoteFS provides API for accessing the files on @@ -136,15 +137,15 @@ func (r *RemoteFS) ReadDir(path string) ([]os.FileInfo, error) { return fileInfos, nil } -func (r *RemoteFS) Open(path string) (File, error) { +func (r *RemoteFS) Open(path string) (sftputils.File, error) { return r.OpenFile(path, os.O_RDONLY) } -func (r *RemoteFS) Create(path string, _ int64) (File, error) { +func (r *RemoteFS) Create(path string, _ int64) (sftputils.File, error) { return r.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC) } -func (r *RemoteFS) OpenFile(path string, flags int) (File, error) { +func (r *RemoteFS) OpenFile(path string, flags int) (sftputils.File, error) { return r.Client.OpenFile(path, flags) } diff --git a/lib/sshutils/sftp/sftp.go b/lib/sshutils/sftp/sftp.go index d0d2f1841e2..bcde8d109ed 100644 --- a/lib/sshutils/sftp/sftp.go +++ b/lib/sshutils/sftp/sftp.go @@ -29,9 +29,7 @@ import ( "net/http" "os" "path" // SFTP requires UNIX-style path separators - "runtime" "strconv" - "strings" "time" "github.com/gravitational/trace" @@ -42,38 +40,7 @@ import ( "github.com/gravitational/teleport/api/observability/tracing/ssh" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/sshutils/scp" -) - -// SFTP request methods. -const ( - // MethodGet opens a file for reading. - MethodGet = "Get" - // MethodPut opens a file for writing. - MethodPut = "Put" - // MethodOpen opens a file. - MethodOpen = "Open" - // MethodSetStat sets a file's stats. - MethodSetStat = "Setstat" - // MethodRename renames a file. - MethodRename = "Rename" - // MethodRmdir removes a directory. - MethodRmdir = "Rmdir" - // MethodMkdir creates a directory. - MethodMkdir = "Mkdir" - // MethodLink creates a hard link. - MethodLink = "Link" - // MethodSymlink creates a symbolic link. - MethodSymlink = "Symlink" - // MethodRemove deletes a file. - MethodRemove = "Remove" - // MethodList lists directory entries. - MethodList = "List" - // MethodStat gets a directory entry's stat info. - MethodStat = "Stat" - // MethodLstat gets a directory entry's stat info, without following symbolic links. - MethodLstat = "Lstat" - // MethodReadlink gets the target of a symbolic link. - MethodReadlink = "Readlink" + "github.com/gravitational/teleport/session/sftputils" ) // FileTransferRequest holds the settings for an SFTP file transfer. @@ -105,8 +72,8 @@ type FileTransferRequest struct { // ModeratedSessionID is the optional ID of a moderated session. ModeratedSessionID string - srcFS FileSystem - dstFS FileSystem + srcFS sftputils.FileSystem + dstFS sftputils.FileSystem } func (req *FileTransferRequest) checkAndSetDefaults() error { @@ -141,67 +108,6 @@ func (req *FileTransferRequest) checkAndSetDefaults() error { return nil } -// File is the file interface required for [FileSystem]. -type File interface { - sftp.WriterAtReaderAt - io.ReadWriteCloser - // Name returns the name of the file. - Name() string - // Stat returns the files stat info. - Stat() (fs.FileInfo, error) -} - -// FileSystem describes file operations to be done either locally or over SFTP. -// -// Note: errors returned by a FileSystem should not be `trace.Wrap()`ed so the -// sftp package can parse os errors. -type FileSystem interface { - io.Closer - // Type returns whether the filesystem is "local" or "remote". - Type() string - // Glob returns matching files of a glob pattern. - Glob(pattern string) ([]string, error) - // Stat returns info about a file. - Stat(path string) (os.FileInfo, error) - // ReadDir returns information about files contained within a directory. - ReadDir(path string) ([]os.FileInfo, error) - // Open opens a file for reading. - Open(path string) (File, error) - // Create creates a new file for writing. - Create(path string, size int64) (File, error) - // Mkdir creates a directory. - Mkdir(path string) error - // Chmod sets file permissions. - Chmod(path string, mode os.FileMode) error - // Chtimes sets file access and modification time. - Chtimes(path string, atime, mtime time.Time) error - // OpenFile opens a file with the given flags. - OpenFile(path string, flags int) (File, error) - // Rename renames a file. - Rename(oldpath, newpath string) error - // Lstat returns info about a file or symlink. - Lstat(name string) (os.FileInfo, error) - // RemoveAll recursively removes a file or directory. - RemoveAll(path string) error - // Link creates a new link. - Link(oldname, newname string) error - // Symlink creates a new symlink. - Symlink(oldname, newname string) error - // Remove removes a file or (empty) directory. - Remove(name string) error - // Chown changes a file's owner and/or group. - Chown(name string, uid, gid int) error - // Truncate truncates a file's contents. - Truncate(name string, size int64) error - // Readlink gets the destination for a symlink. - Readlink(name string) (string, error) - // Getwd gets the current working directory. - Getwd() (string, error) - // RealPath canonicalizes a path name, including resolving ".." and - // following symlinks. - RealPath(path string) (string, error) -} - // HTTPTransferRequest describes file transfer request over HTTP. type HTTPTransferRequest struct { // Src is the source file name @@ -310,14 +216,14 @@ func TransferFiles(ctx context.Context, req *FileTransferRequest) error { return trace.Wrap(err) } for i, srcPath := range req.Sources.Paths { - expandedPath, err := ExpandHomeDir(srcPath) + expandedPath, err := sftputils.ExpandHomeDir(srcPath) if err != nil { return trace.Wrap(err) } req.Sources.Paths[i] = expandedPath } default: - req.srcFS = localFS{} + req.srcFS = sftputils.LocalFS{} } defer req.srcFS.Close() @@ -338,73 +244,19 @@ func TransferFiles(ctx context.Context, req *FileTransferRequest) error { if err != nil { return trace.Wrap(err) } - expandedPath, err := ExpandHomeDir(req.Destination.Path) + expandedPath, err := sftputils.ExpandHomeDir(req.Destination.Path) if err != nil { return trace.Wrap(err) } req.Destination.Path = expandedPath default: - req.dstFS = localFS{} + req.dstFS = sftputils.LocalFS{} } defer req.dstFS.Close() return trace.Wrap(transfer(ctx, req)) } -// PathExpansionError is an [error] indicating that -// path expansion was rejected. -type PathExpansionError struct { - path string -} - -func (p PathExpansionError) Error() string { - return fmt.Sprintf("expanding remote ~user paths is not supported, specify an absolute path instead of %q", p.path) -} - -// ExpandHomeDir evaluates the home directory ('~') in a path. -func ExpandHomeDir(pathStr string) (string, error) { - pfxLen, ok := homeDirPrefixLen(pathStr) - if !ok { - return pathStr, nil - } - - if pfxLen == 1 && len(pathStr) > 1 { - return "", trace.Wrap(PathExpansionError{path: pathStr}) - } - - // if an SFTP path is not absolute, it is assumed to start at the user's - // home directory so just strip the prefix and let the SFTP server - // figure out the correct remote path. - trimmedPath := pathStr[pfxLen:] - // Returning an empty string is supported by SFTP but won't be as clear in - // logs or audit events. Since the SFTP server will be rooted at the user's - // home directory, "." and "" are equivalent in this context. - if trimmedPath == "" { - return ".", nil - } - return trimmedPath, nil -} - -// homeDirPrefixLen returns the length of a set of characters that -// indicates the user wants the path to begin with a user's home -// directory and a bool that indicates whether such a prefix exists. -func homeDirPrefixLen(path string) (int, bool) { - if strings.HasPrefix(path, "~/") { - return 2, true - } - // allow '~\' or '~/' on Windows since '\' is the canonical path - // separator but some users may use '/' instead - if runtime.GOOS == "windows" && strings.HasPrefix(path, `~\`) { - return 2, true - } - - if len(path) >= 1 && path[0] == '~' { - return 1, true - } - - return -1, false -} - // transfer performs file transfers func transfer(ctx context.Context, req *FileTransferRequest) error { // get info of source files and ensure appropriate options were passed @@ -439,7 +291,7 @@ func transfer(ctx context.Context, req *FileTransferRequest) error { if fi.IsDir() && !req.Recursive { // Note: Using an error constructor included in lib/client.IsErrorResolvableWithRelogin, // e.g. BadParameter, will lead to relogin attempt and a completely obscure error message. - return trace.Wrap(&NonRecursiveDirectoryTransferError{Path: match}) + return trace.Wrap(&sftputils.NonRecursiveDirectoryTransferError{Path: match}) } fileInfos = append(fileInfos, fi) } @@ -739,173 +591,3 @@ func newProgressBar(size int64, desc string, writer io.Writer) *unboundedProgres progressbar.OptionSetRenderBlankState(true), )} } - -// NonRecursiveDirectoryTransferError is returned when an attempt is made -// to download a directory without providing the recursive option. -// It's used to distinguish this specific situation in clients which -// do not support the recursive option. -type NonRecursiveDirectoryTransferError struct { - Path string -} - -func (n *NonRecursiveDirectoryTransferError) Error() string { - return fmt.Sprintf("%q is a directory, but the recursive option was not passed", n.Path) -} - -func setstat(req *sftp.Request, fs FileSystem) error { - attrFlags := req.AttrFlags() - attrs := req.Attributes() - - if attrFlags.Acmodtime { - atime := time.Unix(int64(attrs.Atime), 0) - mtime := time.Unix(int64(attrs.Mtime), 0) - - err := fs.Chtimes(req.Filepath, atime, mtime) - if err != nil { - return err - } - } - if attrFlags.Permissions { - err := fs.Chmod(req.Filepath, attrs.FileMode()) - if err != nil { - return err - } - } - if attrFlags.UidGid { - err := fs.Chown(req.Filepath, int(attrs.UID), int(attrs.GID)) - if err != nil { - return err - } - } - if attrFlags.Size { - err := fs.Truncate(req.Filepath, int64(attrs.Size)) - if err != nil { - return err - } - } - - return nil -} - -// HandleFilecmd handles file command requests. If filesys is nil, the local -// filesystem will be used. -func HandleFilecmd(req *sftp.Request, filesys FileSystem) error { - if filesys == nil { - filesys = localFS{} - } - switch req.Method { - case MethodSetStat: - return setstat(req, filesys) - case MethodRename: - if req.Target == "" { - return os.ErrInvalid - } - return filesys.Rename(req.Filepath, req.Target) - case MethodRmdir: - fi, err := filesys.Lstat(req.Filepath) - if err != nil { - return err - } - if !fi.IsDir() { - return fmt.Errorf("%q is not a directory", req.Filepath) - } - return filesys.RemoveAll(req.Filepath) - case MethodMkdir: - return filesys.Mkdir(req.Filepath) - case MethodLink: - if req.Target == "" { - return os.ErrInvalid - } - return filesys.Link(req.Target, req.Filepath) - case MethodSymlink: - if req.Target == "" { - return os.ErrInvalid - } - return filesys.Symlink(req.Target, req.Filepath) - case MethodRemove: - fi, err := filesys.Lstat(req.Filepath) - if err != nil { - return err - } - if fi.IsDir() { - return fmt.Errorf("%q is a directory", req.Filepath) - } - return filesys.Remove(req.Filepath) - default: - return sftp.ErrSSHFxOpUnsupported - } -} - -// listerAt satisfies [sftp.listerAt]. -type listerAt []fs.FileInfo - -func (l listerAt) ListAt(ls []fs.FileInfo, offset int64) (int, error) { - if offset >= int64(len(l)) { - return 0, io.EOF - } - n := copy(ls, l[offset:]) - if n < len(ls) { - return n, io.EOF - } - - return n, nil -} - -// fileName satisfies [fs.FileInfo] but only knows a file's name. This -// is necessary when handling 'readlink' requests in sftpHandler.FileList, -// as only the file's name is known after a readlink call. -type fileName string - -func (f fileName) Name() string { - return string(f) -} - -func (f fileName) Size() int64 { - return 0 -} - -func (f fileName) Mode() fs.FileMode { - return 0 -} - -func (f fileName) ModTime() time.Time { - return time.Time{} -} - -func (f fileName) IsDir() bool { - return false -} - -func (f fileName) Sys() any { - return nil -} - -// HandleFilelist handles file list requests. If filesys is nil, the local -// filesystem will be used. -func HandleFilelist(req *sftp.Request, filesys FileSystem) (sftp.ListerAt, error) { - if filesys == nil { - filesys = localFS{} - } - switch req.Method { - case MethodList: - entries, err := filesys.ReadDir(req.Filepath) - if err != nil { - return nil, err - } - return listerAt(entries), nil - case MethodStat: - fi, err := filesys.Stat(req.Filepath) - if err != nil { - return nil, err - } - return listerAt{fi}, nil - case MethodReadlink: - dst, err := filesys.Readlink(req.Filepath) - if err != nil { - return nil, err - } - return listerAt{fileName(dst)}, nil - default: - return nil, sftp.ErrSSHFxOpUnsupported - } -} diff --git a/lib/sshutils/sftp/sftp_test.go b/lib/sshutils/sftp/sftp_test.go index d0306ab5721..09d22307f9c 100644 --- a/lib/sshutils/sftp/sftp_test.go +++ b/lib/sshutils/sftp/sftp_test.go @@ -24,9 +24,7 @@ import ( cryptorand "crypto/rand" "fmt" "io" - "io/fs" mathrand "math/rand/v2" - "net" "net/http" "net/http/httptest" "os" @@ -34,16 +32,13 @@ import ( "strconv" "strings" "testing" - "time" "github.com/google/go-cmp/cmp" "github.com/gravitational/trace" - "github.com/pkg/sftp" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/utils/log/logtest" + "github.com/gravitational/teleport/session/sftputils" ) const fileMaxSize = 1000 @@ -411,7 +406,7 @@ func TestTransferFiles(t *testing.T) { }, errCheck: func(t require.TestingT, err error, i ...any) { require.EqualError(t, err, fmt.Sprintf(`"%s/src" is a directory, but the recursive option was not passed`, i[0])) - require.ErrorAs(t, err, new(*NonRecursiveDirectoryTransferError)) + require.ErrorAs(t, err, new(*sftputils.NonRecursiveDirectoryTransferError)) }, }, { @@ -466,57 +461,6 @@ func TestTransferFiles(t *testing.T) { } } -func TestHomeDirExpansion(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - path string - expandedPath string - errCheck require.ErrorAssertionFunc - }{ - { - name: "absolute path", - path: "/foo/bar", - expandedPath: "/foo/bar", - }, - { - name: "path with tilde-slash", - path: "~/foo/bar", - expandedPath: "foo/bar", - }, - { - name: "just tilde", - path: "~", - expandedPath: ".", - }, - { - name: "tilde slash", - path: "~/", - expandedPath: ".", - }, - { - name: "~user path", - path: "~user/foo", - errCheck: func(t require.TestingT, err error, i ...any) { - require.ErrorIs(t, err, PathExpansionError{path: "~user/foo"}) - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - expanded, err := ExpandHomeDir(tt.path) - if tt.errCheck == nil { - require.NoError(t, err) - require.Equal(t, tt.expandedPath, expanded) - } else { - tt.errCheck(t, err) - } - }) - } -} - func TestCopyingSymlinkedFile(t *testing.T) { t.Parallel() @@ -543,7 +487,7 @@ func TestCopyingSymlinkedFile(t *testing.T) { } type mockFile struct { - File + sftputils.File altDataSource io.Reader } @@ -552,21 +496,21 @@ func (m *mockFile) Read(p []byte) (int, error) { } type mockFS struct { - localFS + sftputils.LocalFS fileAccesses map[string]int altData io.Reader } -func (m *mockFS) Open(path string) (File, error) { +func (m *mockFS) Open(path string) (sftputils.File, error) { if m.fileAccesses == nil { m.fileAccesses = make(map[string]int) } - realPath, err := m.localFS.RealPath(path) + realPath, err := m.LocalFS.RealPath(path) if err != nil { return nil, trace.Wrap(err) } m.fileAccesses[realPath]++ - file, err := m.localFS.Open(path) + file, err := m.LocalFS.Open(path) if err != nil || m.altData == nil { return file, err } @@ -676,7 +620,7 @@ func TestHTTPUpload(t *testing.T) { }, ) require.NoError(t, err) - transferReq.dstFS = &localFS{} + transferReq.dstFS = &sftputils.LocalFS{} err = TransferFiles(t.Context(), transferReq) require.NoError(t, err) @@ -868,336 +812,3 @@ func compareFileInfos(t *testing.T, preserveAttrs bool, dstInfo, srcInfo os.File // often different when run in CI } } - -type mockCmdHandlers struct { - sftp.Handlers -} - -func (m mockCmdHandlers) Filecmd(req *sftp.Request) error { - return trace.Wrap(HandleFilecmd(req, localFS{})) -} - -func TestHandleFilecmd(t *testing.T) { - t.Parallel() - // We're using a full client/server instead of just calling HandleFilecmd so - // the sftp package can handle marshaling attributes. - clientConn, serverConn := net.Pipe() - srv := sftp.NewRequestServer(serverConn, sftp.Handlers{ - FileGet: sftp.InMemHandler().FileGet, - FilePut: sftp.InMemHandler().FilePut, - FileCmd: mockCmdHandlers{}, - FileList: sftp.InMemHandler().FileList, - }) - - t.Cleanup(func() { require.NoError(t, srv.Close()) }) - go srv.Serve() - - clt, err := sftp.NewClientPipe(clientConn, clientConn) - require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, clt.Close()) }) - - t.Run("chtimes", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "test.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - originalInfo, err := os.Stat(file) - require.NoError(t, err) - setTime := originalInfo.ModTime().Add(time.Hour).Round(time.Second) - - assert.NoError(t, clt.Chtimes(file, setTime, setTime)) - updatedInfo, err := os.Stat(file) - if assert.NoError(t, err) { - assert.Equal(t, setTime, updatedInfo.ModTime()) - } - }) - - t.Run("chmod", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "test.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - - assert.NoError(t, clt.Chmod(file, 0o666)) - fi, err := os.Stat(file) - if assert.NoError(t, err) { - assert.Equal(t, fs.FileMode(0o666), fi.Mode().Perm()) - } - }) - - t.Run("truncate", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "test.txt") - require.NoError(t, os.WriteFile(file, []byte(strings.Repeat("a", 100)), 0o644)) - - assert.NoError(t, clt.Truncate(file, 50)) - data, err := os.ReadFile(file) - if assert.NoError(t, err) { - assert.Len(t, data, 50) - } - }) - - t.Run("rename", func(t *testing.T) { - root := t.TempDir() - initialFile := filepath.Join(root, "foo.txt") - finalFile := filepath.Join(root, "bar.txt") - require.NoError(t, os.WriteFile(initialFile, []byte("test"), 0o644)) - - assert.NoError(t, clt.Rename(initialFile, finalFile)) - assert.NoFileExists(t, initialFile) - assert.FileExists(t, finalFile) - }) - - t.Run("rename missing target", func(t *testing.T) { - root := t.TempDir() - initialFile := filepath.Join(root, "foo.txt") - finalFile := filepath.Join(root, "bar.txt") - assert.Error(t, clt.Rename(initialFile, finalFile)) - assert.NoFileExists(t, finalFile) - }) - - t.Run("rmdir", func(t *testing.T) { - root := t.TempDir() - dir := filepath.Join(root, "foo") - innerFile := filepath.Join(dir, "test.txt") - require.NoError(t, os.Mkdir(dir, defaults.DirectoryPermissions)) - require.NoError(t, os.WriteFile(innerFile, []byte("test"), 0o644)) - - assert.NoError(t, clt.RemoveDirectory(dir)) - assert.NoDirExists(t, dir) - }) - - t.Run("rmdir not found", func(t *testing.T) { - root := t.TempDir() - dir := filepath.Join(root, "foo") - assert.Error(t, clt.RemoveDirectory(dir)) - }) - - t.Run("rmdir not a dir", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "test.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - - assert.Error(t, clt.RemoveDirectory(file)) - assert.FileExists(t, file) - }) - - t.Run("mkdir", func(t *testing.T) { - root := t.TempDir() - outer := filepath.Join(root, "a") - inner := filepath.Join(outer, "b/c") - require.NoError(t, os.Mkdir(outer, defaults.DirectoryPermissions)) - - assert.NoError(t, clt.Mkdir(inner)) - assert.DirExists(t, inner) - }) - - t.Run("link", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "foo.txt") - target := filepath.Join(root, "bar.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - - assert.NoError(t, clt.Link(target, file)) - fi, err := os.Lstat(target) - if assert.NoError(t, err) { - assert.Zero(t, fi.Mode()&os.ModeSymlink) - } - }) - - t.Run("link missing target", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "foo.txt") - target := filepath.Join(root, "bar.txt") - - assert.Error(t, clt.Link(target, file)) - assert.NoFileExists(t, target) - }) - - t.Run("link unset target", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "foo.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - assert.Error(t, clt.Link(file, "")) - }) - - t.Run("symlink", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "foo.txt") - target := filepath.Join(root, "bar.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - - assert.NoError(t, clt.Symlink(target, file)) - fi, err := os.Lstat(target) - assert.NoError(t, err) - assert.NotZero(t, fi.Mode()&os.ModeSymlink) - }) - - t.Run("symlink unset target", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "foo.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - assert.Error(t, clt.Symlink(file, "")) - }) - - t.Run("remove", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "test.txt") - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - - assert.NoError(t, clt.Remove(file)) - assert.NoFileExists(t, file) - }) - - t.Run("remove not found", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "test.txt") - - assert.Error(t, clt.Remove(file)) - }) - - t.Run("remove directory", func(t *testing.T) { - root := t.TempDir() - dir := filepath.Join(root, "dir") - require.NoError(t, os.Mkdir(dir, defaults.DirectoryPermissions)) - - assert.NoError(t, clt.Remove(dir)) - assert.NoDirExists(t, dir) - }) - - t.Run("unsupported operation", func(t *testing.T) { - root := t.TempDir() - file := filepath.Join(root, "test.txt") - require.NoError(t, os.WriteFile(file, []byte("foo"), 0o644)) - req := sftp.NewRequest(MethodStat, file) - assert.Error(t, HandleFilecmd(req, localFS{})) - }) -} - -type fileInfo struct { - name string - mode fs.FileMode - size int64 -} - -func (fi fileInfo) Name() string { - return fi.name -} - -func (fi fileInfo) Size() int64 { - return fi.size -} - -func (fi fileInfo) Mode() fs.FileMode { - return fi.mode -} - -func (fi fileInfo) ModTime() time.Time { - return time.Time{} -} - -func (fi fileInfo) IsDir() bool { - return false -} - -func (fi fileInfo) Sys() any { - return nil -} - -func TestHandleFilelist(t *testing.T) { - t.Parallel() - root := t.TempDir() - statMap := make(map[string]fs.FileInfo, 10) - for i := range 5 { - fileName := fmt.Sprintf("file-%d", i) - file := filepath.Join(root, fileName) - require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) - statMap[fileName] = fileInfo{ - name: fileName, - mode: 0o644, - size: 4, - } - symlinkName := fmt.Sprintf("file-%d", i+5) - symlink := filepath.Join(root, symlinkName) - require.NoError(t, os.Symlink(file, symlink)) - statMap[symlinkName] = fileInfo{ - name: symlinkName, - mode: 0o644, - size: 4, - } - } - - // Add a broken symlink. - brokenSymlinkName := "broken-symlink" - brokenSymlink := filepath.Join(root, brokenSymlinkName) - brokenTarget := filepath.Join(root, "this-file-does-not-exist") - require.NoError(t, os.Symlink(brokenTarget, brokenSymlink)) - symlinkStat, err := os.Lstat(brokenSymlink) - require.NoError(t, err) - statMap[brokenSymlinkName] = fileInfo{ - name: brokenSymlinkName, - mode: symlinkStat.Mode(), - size: int64(len(brokenTarget)), - } - - tests := []struct { - name string - req *sftp.Request - assert assert.ErrorAssertionFunc - expectedOutput map[string]fs.FileInfo - }{ - { - name: "list", - req: sftp.NewRequest(MethodList, root), - assert: assert.NoError, - expectedOutput: statMap, - }, - { - name: "stat", - req: sftp.NewRequest(MethodStat, root+"/file-0"), - assert: assert.NoError, - expectedOutput: map[string]fs.FileInfo{ - "file-0": fileInfo{ - name: "file-0", - mode: 0o644, - size: 4, - }, - }, - }, - { - name: "readlink", - req: sftp.NewRequest(MethodReadlink, root+"/file-5"), - assert: assert.NoError, - expectedOutput: map[string]fs.FileInfo{ - root + "/file-0": fileName(root + "/file-0"), - }, - }, - { - name: "unsupported operation", - req: sftp.NewRequest(MethodRemove, root), - assert: assert.Error, - }, - } - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - lister, err := HandleFilelist(tc.req, localFS{}) - tc.assert(t, err) - if tc.expectedOutput == nil { - assert.Nil(t, lister) - return - } - assert.NotNil(t, lister) - - list := make([]fs.FileInfo, len(tc.expectedOutput)) - n, err := lister.ListAt(list, 0) - assert.NoError(t, err) - assert.Equal(t, len(tc.expectedOutput), n) - for _, fi := range list { - entry, ok := tc.expectedOutput[fi.Name()] - if assert.True(t, ok, "unexpected file %q", fi.Name()) { - assert.Equal(t, entry.Name(), fi.Name()) - assert.Equal(t, entry.Size(), fi.Size(), fi.Name()) - assert.Equal(t, entry.Mode(), fi.Mode(), "%s: expected mode 0o%o, got mode 0o%o", fi.Name(), entry.Mode(), fi.Mode()) - } - } - }) - } -} diff --git a/lib/sshutils/sftp/utils.go b/lib/sshutils/sftp/utils.go index 1c47a55b201..498c9cb9b79 100644 --- a/lib/sshutils/sftp/utils.go +++ b/lib/sshutils/sftp/utils.go @@ -24,26 +24,8 @@ import ( "io" "io/fs" "os" - "sync/atomic" - "time" - - "github.com/gravitational/trace" - "github.com/pkg/sftp" - - apievents "github.com/gravitational/teleport/api/types/events" - "github.com/gravitational/teleport/lib/events" ) -// fileWrapper is a wrapper for *os.File that implements the WriteTo() method -// required for concurrent data transfer. -type fileWrapper struct { - *os.File -} - -func (wt *fileWrapper) WriteTo(w io.Writer) (n int64, err error) { - return io.Copy(w, wt.File) -} - // fileStreamReader is a thin wrapper around fs.File with additional streams. type fileStreamReader struct { ctx context.Context @@ -91,184 +73,3 @@ func (c *cancelWriter) Write(b []byte) (int, error) { } return c.stream.Write(b) } - -// TrackedFile is a [File] that counts the bytes read from/written to it. -type TrackedFile struct { - File - // BytesRead is the number of bytes read. - bytesRead atomic.Uint64 - // BytesWritten is the number of bytes written. - bytesWritten atomic.Uint64 -} - -func (t *TrackedFile) ReadAt(b []byte, off int64) (int, error) { - n, err := t.File.ReadAt(b, off) - t.bytesRead.Add(uint64(n)) - return n, err -} - -func (t *TrackedFile) WriteAt(b []byte, off int64) (int, error) { - n, err := t.File.WriteAt(b, off) - t.bytesWritten.Add(uint64(n)) - return n, err -} - -func (t *TrackedFile) BytesRead() uint64 { - return t.bytesRead.Load() -} - -func (t *TrackedFile) BytesWritten() uint64 { - return t.bytesWritten.Load() -} - -// ParseFlags parses Open flags from an SFTP request to an int as used by -// [os.OpenFile]. -func ParseFlags(req *sftp.Request) int { - pflags := req.Pflags() - var flags int - if pflags.Read && pflags.Write { - flags = os.O_RDWR - } else if pflags.Read { - flags = os.O_RDONLY - } else if pflags.Write { - flags = os.O_WRONLY - } - - if pflags.Append { - flags |= os.O_APPEND - } - if pflags.Creat { - flags |= os.O_CREATE - } - if pflags.Excl { - flags |= os.O_EXCL - } - if pflags.Trunc { - flags |= os.O_TRUNC - } - - return flags -} - -// ParseSFTPEvent parses an SFTP request and associated error into an SFTP -// audit event. -func ParseSFTPEvent(req *sftp.Request, workingDirectory string, reqErr error) (*apievents.SFTP, error) { - event := &apievents.SFTP{ - Metadata: apievents.Metadata{ - Type: events.SFTPEvent, - Time: time.Now(), - }, - } - - switch req.Method { - case MethodOpen, MethodGet, MethodPut: - if reqErr == nil { - event.Code = events.SFTPOpenCode - } else { - event.Code = events.SFTPOpenFailureCode - } - event.Action = apievents.SFTPAction_OPEN - case MethodSetStat: - if reqErr == nil { - event.Code = events.SFTPSetstatCode - } else { - event.Code = events.SFTPSetstatFailureCode - } - event.Action = apievents.SFTPAction_SETSTAT - case MethodList: - if reqErr == nil { - event.Code = events.SFTPReaddirCode - } else { - event.Code = events.SFTPReaddirFailureCode - } - event.Action = apievents.SFTPAction_READDIR - case MethodRemove: - if reqErr == nil { - event.Code = events.SFTPRemoveCode - } else { - event.Code = events.SFTPRemoveFailureCode - } - event.Action = apievents.SFTPAction_REMOVE - case MethodMkdir: - if reqErr == nil { - event.Code = events.SFTPMkdirCode - } else { - event.Code = events.SFTPMkdirFailureCode - } - event.Action = apievents.SFTPAction_MKDIR - case MethodRmdir: - if reqErr == nil { - event.Code = events.SFTPRmdirCode - } else { - event.Code = events.SFTPRmdirFailureCode - } - event.Action = apievents.SFTPAction_RMDIR - case MethodRename: - if reqErr == nil { - event.Code = events.SFTPRenameCode - } else { - event.Code = events.SFTPRenameFailureCode - } - event.Action = apievents.SFTPAction_RENAME - case MethodSymlink: - if reqErr == nil { - event.Code = events.SFTPSymlinkCode - } else { - event.Code = events.SFTPSymlinkFailureCode - } - event.Action = apievents.SFTPAction_SYMLINK - case MethodLink: - if reqErr == nil { - event.Code = events.SFTPLinkCode - } else { - event.Code = events.SFTPLinkFailureCode - } - event.Action = apievents.SFTPAction_LINK - default: - return nil, trace.BadParameter("unknown SFTP request %q", req.Method) - } - - event.Path = req.Filepath - event.TargetPath = req.Target - event.Flags = req.Flags - event.WorkingDirectory = workingDirectory - if req.Method == MethodSetStat { - attrFlags := req.AttrFlags() - attrs := req.Attributes() - event.Attributes = new(apievents.SFTPAttributes) - - if attrFlags.Acmodtime { - atime := time.Unix(int64(attrs.Atime), 0) - mtime := time.Unix(int64(attrs.Mtime), 0) - event.Attributes.AccessTime = &atime - event.Attributes.ModificationTime = &mtime - } - if attrFlags.Permissions { - perms := uint32(attrs.FileMode().Perm()) - event.Attributes.Permissions = &perms - } - if attrFlags.Size { - event.Attributes.FileSize = &attrs.Size - } - if attrFlags.UidGid { - event.Attributes.UID = &attrs.UID - event.Attributes.GID = &attrs.GID - } - } - if reqErr != nil { - // If possible, strip the filename from the error message. The - // path will be included in audit events already, no need to - // make the error message longer than it needs to be. - var pathErr *fs.PathError - var linkErr *os.LinkError - if errors.As(reqErr, &pathErr) { - event.Error = pathErr.Err.Error() - } else if errors.As(reqErr, &linkErr) { - event.Error = linkErr.Err.Error() - } else { - event.Error = reqErr.Error() - } - } - - return event, nil -} diff --git a/lib/tbot/tbot_test.go b/lib/tbot/tbot_test.go index 6a26fc37496..149c3ea1357 100644 --- a/lib/tbot/tbot_test.go +++ b/lib/tbot/tbot_test.go @@ -82,10 +82,12 @@ import ( "github.com/gravitational/teleport/lib/tlsca" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/utils/log/logtest" + "github.com/gravitational/teleport/session/reexec" "github.com/gravitational/teleport/tool/teleport/testenv" ) func TestMain(m *testing.M) { + reexec.MaybeReexec() logtest.InitLogger(testing.Verbose) ctx, cancel := context.WithCancel(context.Background()) diff --git a/lib/teleterm/clusters/cluster_file_transfer.go b/lib/teleterm/clusters/cluster_file_transfer.go index 9ace30c909e..a6a88de075c 100644 --- a/lib/teleterm/clusters/cluster_file_transfer.go +++ b/lib/teleterm/clusters/cluster_file_transfer.go @@ -34,6 +34,7 @@ import ( "github.com/gravitational/teleport/lib/sshutils/sftp" "github.com/gravitational/teleport/lib/teleterm/api/uri" "github.com/gravitational/teleport/lib/utils" + "github.com/gravitational/teleport/session/sftputils" ) type FileTransferProgressSender = func(progress *api.FileTransferProgress) error @@ -81,7 +82,7 @@ func (c *Cluster) TransferFile(ctx context.Context, clt *client.ClusterClient, r err := AddMetadataToRetryableError(ctx, func() error { err := c.clusterClient.TransferFiles(ctx, sftpReq) - if errors.As(err, new(*sftp.NonRecursiveDirectoryTransferError)) { + if errors.As(err, new(*sftputils.NonRecursiveDirectoryTransferError)) { return trace.Errorf("transferring directories through Teleport Connect is not supported at the moment, please use tsh scp -r") } return trace.Wrap(err) diff --git a/lib/web/apiserver_test.go b/lib/web/apiserver_test.go index 6dc8d854a58..6dd8111c8af 100644 --- a/lib/web/apiserver_test.go +++ b/lib/web/apiserver_test.go @@ -189,14 +189,9 @@ type WebSuite struct { // TestMain will re-execute Teleport to run a command if "exec" is passed to // it as an argument. Otherwise, it will run tests as normal. func TestMain(m *testing.M) { + reexec.MaybeReexec() logtest.InitLogger(testing.Verbose) modules.SetInsecureTestMode(true) - // If the test is re-executing itself, execute the command that comes over - // the pipe. - if reexec.IsReexec() { - reexec.RunAndExit(os.Args[1]) - return - } ctx, cancel := context.WithCancel(context.Background()) cryptosuitestest.PrecomputeRSAKeys(ctx) diff --git a/lib/web/files.go b/lib/web/files.go index 38fae09cb69..c6acab7252b 100644 --- a/lib/web/files.go +++ b/lib/web/files.go @@ -45,6 +45,7 @@ import ( "github.com/gravitational/teleport/lib/sshca" "github.com/gravitational/teleport/lib/sshutils/sftp" "github.com/gravitational/teleport/lib/utils" + "github.com/gravitational/teleport/session/sftputils" ) // fileTransferRequest describes HTTP file transfer request @@ -264,7 +265,7 @@ func (h *Handler) transferFile(w http.ResponseWriter, r *http.Request, p httprou } if err := sftp.TransferFiles(ctx, sftpReq); err != nil { - if errors.As(err, new(*sftp.NonRecursiveDirectoryTransferError)) { + if errors.As(err, new(*sftputils.NonRecursiveDirectoryTransferError)) { return nil, trace.Errorf("transferring directories through the Web UI is not supported at the moment, please use tsh scp -r") } diff --git a/session/reexec/exec_test.go b/session/reexec/exec_test.go index ac0867bb9b6..b8e10137bb8 100644 --- a/session/reexec/exec_test.go +++ b/session/reexec/exec_test.go @@ -26,10 +26,7 @@ import ( ) func TestMain(m *testing.M) { - if IsReexec() { - RunAndExit(os.Args[1]) - return - } + MaybeReexec() if !flag.Parsed() { flag.Parse() diff --git a/session/reexec/reexec.go b/session/reexec/reexec.go index c92d95267c4..3711f00ba61 100644 --- a/session/reexec/reexec.go +++ b/session/reexec/reexec.go @@ -56,6 +56,7 @@ import ( "github.com/gravitational/teleport/session/pam/pamcfg" "github.com/gravitational/teleport/session/reexec/internal/logutils" "github.com/gravitational/teleport/session/reexec/reexecconstants" + "github.com/gravitational/teleport/session/reexec/reexecsftp" "github.com/gravitational/teleport/session/selinux" "github.com/gravitational/teleport/session/shell" "github.com/gravitational/teleport/session/uacc" @@ -1122,6 +1123,12 @@ func RunAndExit(commandType string) { code = runCheckHomeDir() case reexecconstants.ParkSubCommand: code = runPark() + case reexecconstants.SFTPSubCommand: + initLogger("sftp", os.Stderr, ExecLogConfig{}) + err = reexecsftp.RunSFTP(slog.Default()) + if err != nil { + code = 1 + } default: code, err = reexecconstants.RemoteCommandFailure, fmt.Errorf("unknown command type: %v", commandType) } @@ -1144,19 +1151,36 @@ func RunAndExit(commandType string) { os.Exit(code) } +// MaybeReexec checks if the command-line arguments are those of a Teleport +// reexec command, and if so, runs the logic for the command (terminating the +// process at the end). Should be the first thing called in the main function +// for the Teleport binary or in the TestMain for packages that rely on +// reexecution. +func MaybeReexec() { + if IsReexec() { + RunAndExit(os.Args[1]) + } +} + +// TODO(espadolini): remove IsReexec and RunAndExit in favor of requiring MaybeReexec, after enterprise is updated + // IsReexec determines if the current process is a teleport reexec command. // Used by tests to reroute the execution to RunAndExit. func IsReexec() bool { - if len(os.Args) >= 2 { - switch os.Args[1] { - case reexecconstants.ExecSubCommand, reexecconstants.NetworkingSubCommand, - reexecconstants.CheckHomeDirSubCommand, - reexecconstants.ParkSubCommand, reexecconstants.SFTPSubCommand: - return true - } + if len(os.Args) < 2 { + return false } - return false + switch os.Args[1] { + case reexecconstants.ExecSubCommand, + reexecconstants.NetworkingSubCommand, + reexecconstants.CheckHomeDirSubCommand, + reexecconstants.ParkSubCommand, + reexecconstants.SFTPSubCommand: + return true + default: + return false + } } // openFileAsUser opens a file as the given user to ensure proper access checks. This is unsafe and should not be used outside of diff --git a/session/reexec/reexecsftp/sess.go b/session/reexec/reexecsftp/sess.go new file mode 100644 index 00000000000..495d813249f --- /dev/null +++ b/session/reexec/reexecsftp/sess.go @@ -0,0 +1,32 @@ +// Teleport +// Copyright (C) 2026 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package reexecsftp + +// FileTransferRequest is a request to upload or download a file from a node. +type FileTransferRequest struct { + // ID is a UUID that uniquely identifies a file transfer request + // and is unlikely to collide with another file transfer request + ID string + // Requester is the Teleport User that requested the file transfer + Requester string + // Download is true if the request is a download, false if its an upload + Download bool + // Filename is the name of the file to upload. + Filename string + // Location of the requested download or where a file will be uploaded + Location string +} diff --git a/tool/teleport/common/sftp.go b/session/reexec/reexecsftp/sftp.go similarity index 82% rename from tool/teleport/common/sftp.go rename to session/reexec/reexecsftp/sftp.go index e7d1a9ea931..82ff9ade21a 100644 --- a/tool/teleport/common/sftp.go +++ b/session/reexec/reexecsftp/sftp.go @@ -16,11 +16,10 @@ * along with this program. If not, see . */ -package common +package reexecsftp import ( "bufio" - "bytes" "context" "encoding/json" "errors" @@ -33,17 +32,11 @@ import ( "sync" "time" - "github.com/gogo/protobuf/jsonpb" //nolint:depguard // needed for backwards compatibility "github.com/gravitational/trace" "github.com/pkg/sftp" "golang.org/x/sys/unix" - "github.com/gravitational/teleport" - apievents "github.com/gravitational/teleport/api/types/events" - "github.com/gravitational/teleport/lib/defaults" - "github.com/gravitational/teleport/lib/events" - "github.com/gravitational/teleport/lib/srv" - sftputils "github.com/gravitational/teleport/lib/sshutils/sftp" + "github.com/gravitational/teleport/session/sftputils" ) type compositeCh struct { @@ -77,10 +70,10 @@ type sftpHandler struct { mtx sync.Mutex files []*sftputils.TrackedFile - events chan<- apievents.AuditEvent + events chan<- sftputils.Event } -func newSFTPHandler(logger *slog.Logger, req *srv.FileTransferRequest, events chan<- apievents.AuditEvent) (*sftpHandler, error) { +func newSFTPHandler(logger *slog.Logger, req *FileTransferRequest, events chan<- sftputils.Event) (*sftpHandler, error) { var allowed *allowedOps if req != nil { allowed = &allowedOps{ @@ -193,7 +186,7 @@ func (s *sftpHandler) openFile(req *sftp.Request) (sftp.WriterAtReaderAt, error) return nil, err } - f, err := os.OpenFile(req.Filepath, sftputils.ParseFlags(req), defaults.FilePermissions) + f, err := os.OpenFile(req.Filepath, sftputils.ParseFlags(req), 0o644) if err != nil { return nil, err } @@ -261,10 +254,10 @@ func (s *sftpHandler) sendSFTPEvent(req *sftp.Request, reqErr error) { } else if reqErr != nil { s.logger.DebugContext(req.Context(), "failed handling SFTP request", "request", req.Method, "error", reqErr) } - s.events <- event + s.events <- sftputils.Event{SFTP: event} } -func onSFTP() error { +func RunSFTP(logger *slog.Logger) error { chr, err := openFD(3, "chr") if err != nil { return trace.Wrap(err) @@ -281,9 +274,6 @@ func onSFTP() error { } defer auditFile.Close() - // Ensure the parent process will receive log messages from us - logger := slog.With(teleport.ComponentKey, teleport.ComponentSubsystemSFTP) - currentUser, err := user.Current() if err != nil { return trace.Wrap(err) @@ -296,7 +286,7 @@ func onSFTP() error { // Read the file transfer request for this session if one exists bufferedReader := bufio.NewReader(chr) var encodedReq []byte - var fileTransferReq *srv.FileTransferRequest + var fileTransferReq *FileTransferRequest for { b, err := bufferedReader.ReadByte() if err != nil { @@ -309,14 +299,14 @@ func onSFTP() error { encodedReq = append(encodedReq, b) } if len(encodedReq) != 0 { - fileTransferReq = new(srv.FileTransferRequest) + fileTransferReq = new(FileTransferRequest) if err := json.Unmarshal(encodedReq, fileTransferReq); err != nil { return trace.Wrap(err) } } ch := compositeCh{io.NopCloser(bufferedReader), chw} - sftpEvents := make(chan apievents.AuditEvent, 1) + sftpEvents := make(chan sftputils.Event, 1) h, err := newSFTPHandler(logger, fileTransferReq, sftpEvents) if err != nil { return trace.Wrap(err) @@ -334,25 +324,10 @@ func onSFTP() error { // process to avoid blocking the SFTP connection on event handling done := make(chan struct{}) go func() { - var m jsonpb.Marshaler - var buf bytes.Buffer + enc := json.NewEncoder(auditFile) + enc.SetEscapeHTML(false) for event := range sftpEvents { - oneOfEvent, err := apievents.ToOneOf(event) - if err != nil { - logger.WarnContext(ctx, "Failed to convert SFTP event to OneOf", "error", err) - continue - } - - buf.Reset() - if err := m.Marshal(&buf, oneOfEvent); err != nil { - logger.WarnContext(ctx, "Failed to marshal SFTP event", "error", err) - continue - } - - // Append a NULL byte so the parent process will know where - // this event ends - buf.WriteByte(0x0) - _, err = io.Copy(auditFile, &buf) + err := enc.Encode(event) if err != nil { logger.WarnContext(ctx, "Failed to send SFTP event to parent", "error", err) } @@ -369,23 +344,20 @@ func onSFTP() error { } // Send a summary event last - summaryEvent := &apievents.SFTPSummary{ - Metadata: apievents.Metadata{ - Type: events.SFTPSummaryEvent, - Code: events.SFTPSummaryCode, - Time: time.Now(), - }, + summaryEvent := &sftputils.SFTPSummaryEvent{ + Time: time.Now().UnixNano(), + Stats: make([]sftputils.SFTPSummaryEventFileTransferStat, 0, len(h.files)), } // We don't need to worry about closing these files, handler will // take care of that for us for _, f := range h.files { - summaryEvent.FileTransferStats = append(summaryEvent.FileTransferStats, &apievents.FileTransferStat{ - Path: f.Name(), - BytesRead: f.BytesRead(), - BytesWritten: f.BytesWritten(), + summaryEvent.Stats = append(summaryEvent.Stats, sftputils.SFTPSummaryEventFileTransferStat{ + Path: f.Name(), + Read: f.BytesRead(), + Written: f.BytesWritten(), }) } - sftpEvents <- summaryEvent + sftpEvents <- sftputils.Event{Summary: summaryEvent} // Wait until event marshaling goroutine is finished close(sftpEvents) diff --git a/lib/sshutils/sftp/local.go b/session/sftputils/local.go similarity index 61% rename from lib/sshutils/sftp/local.go rename to session/sftputils/local.go index fcfc88f9d74..294f3d26754 100644 --- a/lib/sshutils/sftp/local.go +++ b/session/sftputils/local.go @@ -16,34 +16,32 @@ * along with this program. If not, see . */ -package sftp +package sftputils import ( "io/fs" "os" "path/filepath" "time" - - "github.com/gravitational/teleport/lib/defaults" ) -// localFS provides API for accessing the files on +// LocalFS provides API for accessing the files on // the local file system -type localFS struct{} +type LocalFS struct{} -func (l localFS) Type() string { +func (l LocalFS) Type() string { return "local" } -func (l localFS) Glob(pattern string) ([]string, error) { +func (l LocalFS) Glob(pattern string) ([]string, error) { return filepath.Glob(pattern) } -func (l localFS) Stat(path string) (os.FileInfo, error) { +func (l LocalFS) Stat(path string) (os.FileInfo, error) { return os.Stat(path) } -func (l localFS) ReadDir(path string) ([]os.FileInfo, error) { +func (l LocalFS) ReadDir(path string) ([]os.FileInfo, error) { entries, err := os.ReadDir(path) if err != nil { return nil, err @@ -68,7 +66,7 @@ func (l localFS) ReadDir(path string) ([]os.FileInfo, error) { return fileInfos, nil } -func (l localFS) Open(path string) (File, error) { +func (l LocalFS) Open(path string) (File, error) { f, err := os.Open(path) if err != nil { @@ -78,16 +76,16 @@ func (l localFS) Open(path string) (File, error) { return &fileWrapper{File: f}, nil } -func (l localFS) Create(path string, _ int64) (File, error) { +func (l LocalFS) Create(path string, _ int64) (File, error) { return l.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC) } -func (l localFS) OpenFile(path string, flags int) (File, error) { - return os.OpenFile(path, flags, defaults.FilePermissions) +func (l LocalFS) OpenFile(path string, flags int) (File, error) { + return os.OpenFile(path, flags, 0o644) } -func (l localFS) Mkdir(path string) error { - err := os.MkdirAll(path, defaults.DirectoryPermissions) +func (l LocalFS) Mkdir(path string) error { + err := os.MkdirAll(path, 0o755) if err != nil && !os.IsExist(err) { return err } @@ -95,59 +93,59 @@ func (l localFS) Mkdir(path string) error { return nil } -func (l localFS) Chmod(path string, mode os.FileMode) error { +func (l LocalFS) Chmod(path string, mode os.FileMode) error { return os.Chmod(path, mode) } -func (l localFS) Chtimes(path string, atime, mtime time.Time) error { +func (l LocalFS) Chtimes(path string, atime, mtime time.Time) error { return os.Chtimes(path, atime, mtime) } -func (l localFS) Rename(oldpath, newpath string) error { +func (l LocalFS) Rename(oldpath, newpath string) error { return os.Rename(oldpath, newpath) } -func (l localFS) Lstat(name string) (os.FileInfo, error) { +func (l LocalFS) Lstat(name string) (os.FileInfo, error) { return os.Lstat(name) } -func (l localFS) RemoveAll(path string) error { +func (l LocalFS) RemoveAll(path string) error { return os.RemoveAll(path) } -func (l localFS) Link(oldname, newname string) error { +func (l LocalFS) Link(oldname, newname string) error { return os.Link(oldname, newname) } -func (l localFS) Symlink(oldname, newname string) error { +func (l LocalFS) Symlink(oldname, newname string) error { return os.Symlink(oldname, newname) } -func (l localFS) Remove(name string) error { +func (l LocalFS) Remove(name string) error { return os.Remove(name) } -func (l localFS) Chown(name string, uid, gid int) error { +func (l LocalFS) Chown(name string, uid, gid int) error { return os.Chown(name, uid, gid) } -func (l localFS) Truncate(name string, size int64) error { +func (l LocalFS) Truncate(name string, size int64) error { return os.Truncate(name, size) } -func (l localFS) Readlink(name string) (string, error) { +func (l LocalFS) Readlink(name string) (string, error) { return os.Readlink(name) } -func (l localFS) Getwd() (string, error) { +func (l LocalFS) Getwd() (string, error) { return os.Getwd() } -func (l localFS) RealPath(path string) (string, error) { +func (l LocalFS) RealPath(path string) (string, error) { return Realpath(path) } -func (l localFS) Close() error { +func (l LocalFS) Close() error { return nil } diff --git a/lib/sshutils/sftp/local_test.go b/session/sftputils/local_test.go similarity index 99% rename from lib/sshutils/sftp/local_test.go rename to session/sftputils/local_test.go index 3735c6deab0..7c7069021a3 100644 --- a/lib/sshutils/sftp/local_test.go +++ b/session/sftputils/local_test.go @@ -14,7 +14,7 @@ // You should have received a copy of the GNU Affero General Public License // along with this program. If not, see . -package sftp +package sftputils import ( "os" diff --git a/session/sftputils/sftp.go b/session/sftputils/sftp.go new file mode 100644 index 00000000000..aa8568cbd22 --- /dev/null +++ b/session/sftputils/sftp.go @@ -0,0 +1,349 @@ +/* + * Teleport + * Copyright (C) 2023 Gravitational, Inc. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +package sftputils + +import ( + "fmt" + "io" + "io/fs" + "os" + "runtime" + "strings" + "time" + + "github.com/gravitational/trace" + "github.com/pkg/sftp" +) + +// SFTP request methods. +const ( + // MethodGet opens a file for reading. + MethodGet = "Get" + // MethodPut opens a file for writing. + MethodPut = "Put" + // MethodOpen opens a file. + MethodOpen = "Open" + // MethodSetStat sets a file's stats. + MethodSetStat = "Setstat" + // MethodRename renames a file. + MethodRename = "Rename" + // MethodRmdir removes a directory. + MethodRmdir = "Rmdir" + // MethodMkdir creates a directory. + MethodMkdir = "Mkdir" + // MethodLink creates a hard link. + MethodLink = "Link" + // MethodSymlink creates a symbolic link. + MethodSymlink = "Symlink" + // MethodRemove deletes a file. + MethodRemove = "Remove" + // MethodList lists directory entries. + MethodList = "List" + // MethodStat gets a directory entry's stat info. + MethodStat = "Stat" + // MethodLstat gets a directory entry's stat info, without following symbolic links. + MethodLstat = "Lstat" + // MethodReadlink gets the target of a symbolic link. + MethodReadlink = "Readlink" +) + +// File is the file interface required for [FileSystem]. +type File interface { + sftp.WriterAtReaderAt + io.ReadWriteCloser + // Name returns the name of the file. + Name() string + // Stat returns the files stat info. + Stat() (fs.FileInfo, error) +} + +// FileSystem describes file operations to be done either locally or over SFTP. +// +// Note: errors returned by a FileSystem should not be `trace.Wrap()`ed so the +// sftp package can parse os errors. +type FileSystem interface { + io.Closer + // Type returns whether the filesystem is "local" or "remote". + Type() string + // Glob returns matching files of a glob pattern. + Glob(pattern string) ([]string, error) + // Stat returns info about a file. + Stat(path string) (os.FileInfo, error) + // ReadDir returns information about files contained within a directory. + ReadDir(path string) ([]os.FileInfo, error) + // Open opens a file for reading. + Open(path string) (File, error) + // Create creates a new file for writing. + Create(path string, size int64) (File, error) + // Mkdir creates a directory. + Mkdir(path string) error + // Chmod sets file permissions. + Chmod(path string, mode os.FileMode) error + // Chtimes sets file access and modification time. + Chtimes(path string, atime, mtime time.Time) error + // OpenFile opens a file with the given flags. + OpenFile(path string, flags int) (File, error) + // Rename renames a file. + Rename(oldpath, newpath string) error + // Lstat returns info about a file or symlink. + Lstat(name string) (os.FileInfo, error) + // RemoveAll recursively removes a file or directory. + RemoveAll(path string) error + // Link creates a new link. + Link(oldname, newname string) error + // Symlink creates a new symlink. + Symlink(oldname, newname string) error + // Remove removes a file or (empty) directory. + Remove(name string) error + // Chown changes a file's owner and/or group. + Chown(name string, uid, gid int) error + // Truncate truncates a file's contents. + Truncate(name string, size int64) error + // Readlink gets the destination for a symlink. + Readlink(name string) (string, error) + // Getwd gets the current working directory. + Getwd() (string, error) + // RealPath canonicalizes a path name, including resolving ".." and + // following symlinks. + RealPath(path string) (string, error) +} + +// PathExpansionError is an [error] indicating that +// path expansion was rejected. +type PathExpansionError struct { + path string +} + +func (p PathExpansionError) Error() string { + return fmt.Sprintf("expanding remote ~user paths is not supported, specify an absolute path instead of %q", p.path) +} + +// ExpandHomeDir evaluates the home directory ('~') in a path. +func ExpandHomeDir(pathStr string) (string, error) { + pfxLen, ok := homeDirPrefixLen(pathStr) + if !ok { + return pathStr, nil + } + + if pfxLen == 1 && len(pathStr) > 1 { + return "", trace.Wrap(PathExpansionError{path: pathStr}) + } + + // if an SFTP path is not absolute, it is assumed to start at the user's + // home directory so just strip the prefix and let the SFTP server + // figure out the correct remote path. + trimmedPath := pathStr[pfxLen:] + // Returning an empty string is supported by SFTP but won't be as clear in + // logs or audit events. Since the SFTP server will be rooted at the user's + // home directory, "." and "" are equivalent in this context. + if trimmedPath == "" { + return ".", nil + } + return trimmedPath, nil +} + +// homeDirPrefixLen returns the length of a set of characters that +// indicates the user wants the path to begin with a user's home +// directory and a bool that indicates whether such a prefix exists. +func homeDirPrefixLen(path string) (int, bool) { + if strings.HasPrefix(path, "~/") { + return 2, true + } + // allow '~\' or '~/' on Windows since '\' is the canonical path + // separator but some users may use '/' instead + if runtime.GOOS == "windows" && strings.HasPrefix(path, `~\`) { + return 2, true + } + + if len(path) >= 1 && path[0] == '~' { + return 1, true + } + + return -1, false +} + +// NonRecursiveDirectoryTransferError is returned when an attempt is made +// to download a directory without providing the recursive option. +// It's used to distinguish this specific situation in clients which +// do not support the recursive option. +type NonRecursiveDirectoryTransferError struct { + Path string +} + +func (n *NonRecursiveDirectoryTransferError) Error() string { + return fmt.Sprintf("%q is a directory, but the recursive option was not passed", n.Path) +} + +func setstat(req *sftp.Request, fs FileSystem) error { + attrFlags := req.AttrFlags() + attrs := req.Attributes() + + if attrFlags.Acmodtime { + atime := time.Unix(int64(attrs.Atime), 0) + mtime := time.Unix(int64(attrs.Mtime), 0) + + err := fs.Chtimes(req.Filepath, atime, mtime) + if err != nil { + return err + } + } + if attrFlags.Permissions { + err := fs.Chmod(req.Filepath, attrs.FileMode()) + if err != nil { + return err + } + } + if attrFlags.UidGid { + err := fs.Chown(req.Filepath, int(attrs.UID), int(attrs.GID)) + if err != nil { + return err + } + } + if attrFlags.Size { + err := fs.Truncate(req.Filepath, int64(attrs.Size)) + if err != nil { + return err + } + } + + return nil +} + +// HandleFilecmd handles file command requests. If filesys is nil, the local +// filesystem will be used. +func HandleFilecmd(req *sftp.Request, filesys FileSystem) error { + if filesys == nil { + filesys = LocalFS{} + } + switch req.Method { + case MethodSetStat: + return setstat(req, filesys) + case MethodRename: + if req.Target == "" { + return os.ErrInvalid + } + return filesys.Rename(req.Filepath, req.Target) + case MethodRmdir: + fi, err := filesys.Lstat(req.Filepath) + if err != nil { + return err + } + if !fi.IsDir() { + return fmt.Errorf("%q is not a directory", req.Filepath) + } + return filesys.RemoveAll(req.Filepath) + case MethodMkdir: + return filesys.Mkdir(req.Filepath) + case MethodLink: + if req.Target == "" { + return os.ErrInvalid + } + return filesys.Link(req.Target, req.Filepath) + case MethodSymlink: + if req.Target == "" { + return os.ErrInvalid + } + return filesys.Symlink(req.Target, req.Filepath) + case MethodRemove: + fi, err := filesys.Lstat(req.Filepath) + if err != nil { + return err + } + if fi.IsDir() { + return fmt.Errorf("%q is a directory", req.Filepath) + } + return filesys.Remove(req.Filepath) + default: + return sftp.ErrSSHFxOpUnsupported + } +} + +// listerAt satisfies [sftp.listerAt]. +type listerAt []fs.FileInfo + +func (l listerAt) ListAt(ls []fs.FileInfo, offset int64) (int, error) { + if offset >= int64(len(l)) { + return 0, io.EOF + } + n := copy(ls, l[offset:]) + if n < len(ls) { + return n, io.EOF + } + + return n, nil +} + +// fileName satisfies [fs.FileInfo] but only knows a file's name. This +// is necessary when handling 'readlink' requests in sftpHandler.FileList, +// as only the file's name is known after a readlink call. +type fileName string + +func (f fileName) Name() string { + return string(f) +} + +func (f fileName) Size() int64 { + return 0 +} + +func (f fileName) Mode() fs.FileMode { + return 0 +} + +func (f fileName) ModTime() time.Time { + return time.Time{} +} + +func (f fileName) IsDir() bool { + return false +} + +func (f fileName) Sys() any { + return nil +} + +// HandleFilelist handles file list requests. If filesys is nil, the local +// filesystem will be used. +func HandleFilelist(req *sftp.Request, filesys FileSystem) (sftp.ListerAt, error) { + if filesys == nil { + filesys = LocalFS{} + } + switch req.Method { + case MethodList: + entries, err := filesys.ReadDir(req.Filepath) + if err != nil { + return nil, err + } + return listerAt(entries), nil + case MethodStat: + fi, err := filesys.Stat(req.Filepath) + if err != nil { + return nil, err + } + return listerAt{fi}, nil + case MethodReadlink: + dst, err := filesys.Readlink(req.Filepath) + if err != nil { + return nil, err + } + return listerAt{fileName(dst)}, nil + default: + return nil, sftp.ErrSSHFxOpUnsupported + } +} diff --git a/session/sftputils/sftp_test.go b/session/sftputils/sftp_test.go new file mode 100644 index 00000000000..22a3f6d58fd --- /dev/null +++ b/session/sftputils/sftp_test.go @@ -0,0 +1,419 @@ +/* + * Teleport + * Copyright (C) 2023 Gravitational, Inc. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +package sftputils + +import ( + "fmt" + "io/fs" + "net" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/gravitational/trace" + "github.com/pkg/sftp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestHomeDirExpansion(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + path string + expandedPath string + errCheck require.ErrorAssertionFunc + }{ + { + name: "absolute path", + path: "/foo/bar", + expandedPath: "/foo/bar", + }, + { + name: "path with tilde-slash", + path: "~/foo/bar", + expandedPath: "foo/bar", + }, + { + name: "just tilde", + path: "~", + expandedPath: ".", + }, + { + name: "tilde slash", + path: "~/", + expandedPath: ".", + }, + { + name: "~user path", + path: "~user/foo", + errCheck: func(t require.TestingT, err error, i ...any) { + require.ErrorIs(t, err, PathExpansionError{path: "~user/foo"}) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + expanded, err := ExpandHomeDir(tt.path) + if tt.errCheck == nil { + require.NoError(t, err) + require.Equal(t, tt.expandedPath, expanded) + } else { + tt.errCheck(t, err) + } + }) + } +} + +type mockCmdHandlers struct { + sftp.Handlers +} + +func (m mockCmdHandlers) Filecmd(req *sftp.Request) error { + return trace.Wrap(HandleFilecmd(req, LocalFS{})) +} + +func TestHandleFilecmd(t *testing.T) { + t.Parallel() + // We're using a full client/server instead of just calling HandleFilecmd so + // the sftp package can handle marshaling attributes. + clientConn, serverConn := net.Pipe() + srv := sftp.NewRequestServer(serverConn, sftp.Handlers{ + FileGet: sftp.InMemHandler().FileGet, + FilePut: sftp.InMemHandler().FilePut, + FileCmd: mockCmdHandlers{}, + FileList: sftp.InMemHandler().FileList, + }) + + t.Cleanup(func() { require.NoError(t, srv.Close()) }) + go srv.Serve() + + clt, err := sftp.NewClientPipe(clientConn, clientConn) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, clt.Close()) }) + + t.Run("chtimes", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "test.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + originalInfo, err := os.Stat(file) + require.NoError(t, err) + setTime := originalInfo.ModTime().Add(time.Hour).Round(time.Second) + + assert.NoError(t, clt.Chtimes(file, setTime, setTime)) + updatedInfo, err := os.Stat(file) + if assert.NoError(t, err) { + assert.Equal(t, setTime, updatedInfo.ModTime()) + } + }) + + t.Run("chmod", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "test.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + + assert.NoError(t, clt.Chmod(file, 0o666)) + fi, err := os.Stat(file) + if assert.NoError(t, err) { + assert.Equal(t, fs.FileMode(0o666), fi.Mode().Perm()) + } + }) + + t.Run("truncate", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "test.txt") + require.NoError(t, os.WriteFile(file, []byte(strings.Repeat("a", 100)), 0o644)) + + assert.NoError(t, clt.Truncate(file, 50)) + data, err := os.ReadFile(file) + if assert.NoError(t, err) { + assert.Len(t, data, 50) + } + }) + + t.Run("rename", func(t *testing.T) { + root := t.TempDir() + initialFile := filepath.Join(root, "foo.txt") + finalFile := filepath.Join(root, "bar.txt") + require.NoError(t, os.WriteFile(initialFile, []byte("test"), 0o644)) + + assert.NoError(t, clt.Rename(initialFile, finalFile)) + assert.NoFileExists(t, initialFile) + assert.FileExists(t, finalFile) + }) + + t.Run("rename missing target", func(t *testing.T) { + root := t.TempDir() + initialFile := filepath.Join(root, "foo.txt") + finalFile := filepath.Join(root, "bar.txt") + assert.Error(t, clt.Rename(initialFile, finalFile)) + assert.NoFileExists(t, finalFile) + }) + + t.Run("rmdir", func(t *testing.T) { + root := t.TempDir() + dir := filepath.Join(root, "foo") + innerFile := filepath.Join(dir, "test.txt") + require.NoError(t, os.Mkdir(dir, 0o755)) + require.NoError(t, os.WriteFile(innerFile, []byte("test"), 0o644)) + + assert.NoError(t, clt.RemoveDirectory(dir)) + assert.NoDirExists(t, dir) + }) + + t.Run("rmdir not found", func(t *testing.T) { + root := t.TempDir() + dir := filepath.Join(root, "foo") + assert.Error(t, clt.RemoveDirectory(dir)) + }) + + t.Run("rmdir not a dir", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "test.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + + assert.Error(t, clt.RemoveDirectory(file)) + assert.FileExists(t, file) + }) + + t.Run("mkdir", func(t *testing.T) { + root := t.TempDir() + outer := filepath.Join(root, "a") + inner := filepath.Join(outer, "b/c") + require.NoError(t, os.Mkdir(outer, 0o755)) + + assert.NoError(t, clt.Mkdir(inner)) + assert.DirExists(t, inner) + }) + + t.Run("link", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "foo.txt") + target := filepath.Join(root, "bar.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + + assert.NoError(t, clt.Link(target, file)) + fi, err := os.Lstat(target) + if assert.NoError(t, err) { + assert.Zero(t, fi.Mode()&os.ModeSymlink) + } + }) + + t.Run("link missing target", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "foo.txt") + target := filepath.Join(root, "bar.txt") + + assert.Error(t, clt.Link(target, file)) + assert.NoFileExists(t, target) + }) + + t.Run("link unset target", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "foo.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + assert.Error(t, clt.Link(file, "")) + }) + + t.Run("symlink", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "foo.txt") + target := filepath.Join(root, "bar.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + + assert.NoError(t, clt.Symlink(target, file)) + fi, err := os.Lstat(target) + assert.NoError(t, err) + assert.NotZero(t, fi.Mode()&os.ModeSymlink) + }) + + t.Run("symlink unset target", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "foo.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + assert.Error(t, clt.Symlink(file, "")) + }) + + t.Run("remove", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "test.txt") + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + + assert.NoError(t, clt.Remove(file)) + assert.NoFileExists(t, file) + }) + + t.Run("remove not found", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "test.txt") + + assert.Error(t, clt.Remove(file)) + }) + + t.Run("remove directory", func(t *testing.T) { + root := t.TempDir() + dir := filepath.Join(root, "dir") + require.NoError(t, os.Mkdir(dir, 0o755)) + + assert.NoError(t, clt.Remove(dir)) + assert.NoDirExists(t, dir) + }) + + t.Run("unsupported operation", func(t *testing.T) { + root := t.TempDir() + file := filepath.Join(root, "test.txt") + require.NoError(t, os.WriteFile(file, []byte("foo"), 0o644)) + req := sftp.NewRequest(MethodStat, file) + assert.Error(t, HandleFilecmd(req, LocalFS{})) + }) +} + +type fileInfo struct { + name string + mode fs.FileMode + size int64 +} + +func (fi fileInfo) Name() string { + return fi.name +} + +func (fi fileInfo) Size() int64 { + return fi.size +} + +func (fi fileInfo) Mode() fs.FileMode { + return fi.mode +} + +func (fi fileInfo) ModTime() time.Time { + return time.Time{} +} + +func (fi fileInfo) IsDir() bool { + return false +} + +func (fi fileInfo) Sys() any { + return nil +} + +func TestHandleFilelist(t *testing.T) { + t.Parallel() + root := t.TempDir() + statMap := make(map[string]fs.FileInfo, 10) + for i := range 5 { + fileName := fmt.Sprintf("file-%d", i) + file := filepath.Join(root, fileName) + require.NoError(t, os.WriteFile(file, []byte("test"), 0o644)) + statMap[fileName] = fileInfo{ + name: fileName, + mode: 0o644, + size: 4, + } + symlinkName := fmt.Sprintf("file-%d", i+5) + symlink := filepath.Join(root, symlinkName) + require.NoError(t, os.Symlink(file, symlink)) + statMap[symlinkName] = fileInfo{ + name: symlinkName, + mode: 0o644, + size: 4, + } + } + + // Add a broken symlink. + brokenSymlinkName := "broken-symlink" + brokenSymlink := filepath.Join(root, brokenSymlinkName) + brokenTarget := filepath.Join(root, "this-file-does-not-exist") + require.NoError(t, os.Symlink(brokenTarget, brokenSymlink)) + symlinkStat, err := os.Lstat(brokenSymlink) + require.NoError(t, err) + statMap[brokenSymlinkName] = fileInfo{ + name: brokenSymlinkName, + mode: symlinkStat.Mode(), + size: int64(len(brokenTarget)), + } + + tests := []struct { + name string + req *sftp.Request + assert assert.ErrorAssertionFunc + expectedOutput map[string]fs.FileInfo + }{ + { + name: "list", + req: sftp.NewRequest(MethodList, root), + assert: assert.NoError, + expectedOutput: statMap, + }, + { + name: "stat", + req: sftp.NewRequest(MethodStat, root+"/file-0"), + assert: assert.NoError, + expectedOutput: map[string]fs.FileInfo{ + "file-0": fileInfo{ + name: "file-0", + mode: 0o644, + size: 4, + }, + }, + }, + { + name: "readlink", + req: sftp.NewRequest(MethodReadlink, root+"/file-5"), + assert: assert.NoError, + expectedOutput: map[string]fs.FileInfo{ + root + "/file-0": fileName(root + "/file-0"), + }, + }, + { + name: "unsupported operation", + req: sftp.NewRequest(MethodRemove, root), + assert: assert.Error, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + lister, err := HandleFilelist(tc.req, LocalFS{}) + tc.assert(t, err) + if tc.expectedOutput == nil { + assert.Nil(t, lister) + return + } + assert.NotNil(t, lister) + + list := make([]fs.FileInfo, len(tc.expectedOutput)) + n, err := lister.ListAt(list, 0) + assert.NoError(t, err) + assert.Equal(t, len(tc.expectedOutput), n) + for _, fi := range list { + entry, ok := tc.expectedOutput[fi.Name()] + if assert.True(t, ok, "unexpected file %q", fi.Name()) { + assert.Equal(t, entry.Name(), fi.Name()) + assert.Equal(t, entry.Size(), fi.Size(), fi.Name()) + assert.Equal(t, entry.Mode(), fi.Mode(), "%s: expected mode 0o%o, got mode 0o%o", fi.Name(), entry.Mode(), fi.Mode()) + } + } + }) + } +} diff --git a/session/sftputils/utils.go b/session/sftputils/utils.go new file mode 100644 index 00000000000..a1da47d86ef --- /dev/null +++ b/session/sftputils/utils.go @@ -0,0 +1,233 @@ +/* + * Teleport + * Copyright (C) 2023 Gravitational, Inc. + * + * This program is free software: you can redistribute it and/or modify + * it under the terms of the GNU Affero General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * This program is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU Affero General Public License for more details. + * + * You should have received a copy of the GNU Affero General Public License + * along with this program. If not, see . + */ + +package sftputils + +import ( + "errors" + "io" + "io/fs" + "os" + "sync/atomic" + "time" + + "github.com/gravitational/trace" + "github.com/pkg/sftp" +) + +// fileWrapper is a wrapper for *os.File that implements the WriteTo() method +// required for concurrent data transfer. +type fileWrapper struct { + *os.File +} + +func (wt *fileWrapper) WriteTo(w io.Writer) (n int64, err error) { + return io.Copy(w, wt.File) +} + +// TrackedFile is a [File] that counts the bytes read from/written to it. +type TrackedFile struct { + File + // BytesRead is the number of bytes read. + bytesRead atomic.Uint64 + // BytesWritten is the number of bytes written. + bytesWritten atomic.Uint64 +} + +func (t *TrackedFile) ReadAt(b []byte, off int64) (int, error) { + n, err := t.File.ReadAt(b, off) + t.bytesRead.Add(uint64(n)) + return n, err +} + +func (t *TrackedFile) WriteAt(b []byte, off int64) (int, error) { + n, err := t.File.WriteAt(b, off) + t.bytesWritten.Add(uint64(n)) + return n, err +} + +func (t *TrackedFile) BytesRead() uint64 { + return t.bytesRead.Load() +} + +func (t *TrackedFile) BytesWritten() uint64 { + return t.bytesWritten.Load() +} + +// ParseFlags parses Open flags from an SFTP request to an int as used by +// [os.OpenFile]. +func ParseFlags(req *sftp.Request) int { + pflags := req.Pflags() + var flags int + if pflags.Read && pflags.Write { + flags = os.O_RDWR + } else if pflags.Read { + flags = os.O_RDONLY + } else if pflags.Write { + flags = os.O_WRONLY + } + + if pflags.Append { + flags |= os.O_APPEND + } + if pflags.Creat { + flags |= os.O_CREATE + } + if pflags.Excl { + flags |= os.O_EXCL + } + if pflags.Trunc { + flags |= os.O_TRUNC + } + + return flags +} + +// Event is an audit log event passed from the SFTP server process back to the +// main Teleport process. Only one field at a time should be set. +type Event struct { + SFTP *SFTPEvent `json:",omitempty"` + Summary *SFTPSummaryEvent `json:",omitempty"` +} + +// SFTPEvent is an event generated in response to a SFTP operation. +type SFTPEvent struct { + // Time is the event timestamp in nanos since the Unix epoch. + Time int64 + // Method is the SFTP method. + Method string + // Error, if non-empty, signifies that the request has failed and Error + // contains the error message. + Error string `json:",omitempty"` + // Path is the filepath sent by the client. + Path string + // Target is the new path sent by the client for rename and link operations. + Target string `json:",omitempty"` + // Flags are the numerical SFTP flags for the operations, with meaning that + // depends on the operation (Open/Write or SetStat, typically). + Flags uint32 + // WorkDir is the directory that the SFTP server process is in. + WorkDir string + // Attrs contains attributes, currently only populated for SetStat. + Attrs *SFTPEventAttributes `json:",omitempty"` +} + +type SFTPEventAttributes struct { + // Atime is the file access time in seconds since the Unix epoch. + Atime *uint32 `json:",omitempty"` + // Mtime is the file modification time in seconds since the Unix epoch. + Mtime *uint32 `json:",omitempty"` + // Perms is the file permissions. + Perms *uint32 `json:",omitempty"` + // Size is the file size. + Size *uint64 `json:",omitempty"` + // UID is the numerical owner of the file. + UID *uint32 `json:",omitempty"` + // GID is the numerical group of the file. + GID *uint32 `json:",omitempty"` +} + +// SFTPSummaryEvent is generated before the SFTP server process exits. +type SFTPSummaryEvent struct { + // Time is the event timestamp in nanos since the Unix epoch. + Time int64 + // Stats is stats for files that this SFTP session has interacted with. + Stats []SFTPSummaryEventFileTransferStat `json:",omitempty"` +} + +type SFTPSummaryEventFileTransferStat struct { + // Path is the path of the file. + Path string + // Read is the count of bytes read. + Read uint64 + // Written is the count of bytes written. + Written uint64 +} + +// ParseSFTPEvent parses an SFTP request and associated error into an SFTP audit +// event. Changes to this function should be reflected in +// [sshutils/sftp.SFTPEventToProto]. +func ParseSFTPEvent(req *sftp.Request, workingDirectory string, reqErr error) (*SFTPEvent, error) { + event := &SFTPEvent{ + Time: time.Now().UnixNano(), + } + + switch req.Method { + case MethodOpen, MethodGet, MethodPut: + case MethodSetStat: + case MethodList: + case MethodRemove: + case MethodMkdir: + case MethodRmdir: + case MethodRename: + case MethodSymlink: + case MethodLink: + default: + return nil, trace.BadParameter("unknown SFTP request %+q", req.Method) + } + + event.Method = req.Method + event.Path = req.Filepath + event.Target = req.Target + event.Flags = req.Flags + event.WorkDir = workingDirectory + + if req.Method == MethodSetStat { + attrFlags := req.AttrFlags() + attrs := *req.Attributes() + event.Attrs = new(SFTPEventAttributes) + + if attrFlags.Acmodtime { + event.Attrs.Atime = &attrs.Atime + event.Attrs.Mtime = &attrs.Mtime + } + if attrFlags.Permissions { + perms := uint32(attrs.FileMode().Perm()) + event.Attrs.Perms = &perms + } + if attrFlags.Size { + event.Attrs.Size = &attrs.Size + } + if attrFlags.UidGid { + event.Attrs.UID = &attrs.UID + event.Attrs.GID = &attrs.GID + } + } + + if reqErr != nil { + // If possible, strip the filename from the error message. The + // path will be included in audit events already, no need to + // make the error message longer than it needs to be. + var pathErr *fs.PathError + var linkErr *os.LinkError + if errors.As(reqErr, &pathErr) { + event.Error = pathErr.Err.Error() + } else if errors.As(reqErr, &linkErr) { + event.Error = linkErr.Err.Error() + } else { + event.Error = reqErr.Error() + } + if event.Error == "" { + // we signal the failure of a request by the presence of an error + // string, so it must be nonempty here + event.Error = "SFTP request failed with no error message" + } + } + + return event, nil +} diff --git a/tool/teleport/common/teleport.go b/tool/teleport/common/teleport.go index a47e4ecc2da..d16c71b16e3 100644 --- a/tool/teleport/common/teleport.go +++ b/tool/teleport/common/teleport.go @@ -738,8 +738,6 @@ Examples: } case scpc.FullCommand(): err = onSCP(&scpFlags) - case sftp.FullCommand(): - err = onSFTP() case status.FullCommand(): err = onStatus() case dump.FullCommand(): @@ -747,6 +745,11 @@ Examples: case dumpNodeConfigure.FullCommand(): dumpFlags.Roles = defaults.RoleNode err = onConfigDump(dumpFlags) + + // TODO(espadolini): replace these after enterprise calls reexec.MaybeReexec + // in main with an error message ("invalid format for reexec subcommand") + // because if we got here it's because MaybeReexec didn't find the correct + // first argument case exec.FullCommand(): reexec.RunAndExit(reexecconstants.ExecSubCommand) case networking.FullCommand(): @@ -755,6 +758,9 @@ Examples: reexec.RunAndExit(reexecconstants.CheckHomeDirSubCommand) case park.FullCommand(): reexec.RunAndExit(reexecconstants.ParkSubCommand) + case sftp.FullCommand(): + reexec.RunAndExit(reexecconstants.SFTPSubCommand) + case waitNoResolveCmd.FullCommand(): err = onWaitNoResolve(waitFlags) case waitDurationCmd.FullCommand(): diff --git a/tool/teleport/main.go b/tool/teleport/main.go index bb5acfa210a..c6c83cce385 100644 --- a/tool/teleport/main.go +++ b/tool/teleport/main.go @@ -22,6 +22,7 @@ import ( "os" "github.com/gravitational/teleport/lib/observability/metrics" + "github.com/gravitational/teleport/session/reexec" "github.com/gravitational/teleport/tool/teleport/common" ) @@ -30,6 +31,8 @@ func init() { } func main() { + reexec.MaybeReexec() + common.Run(common.Options{ Args: os.Args[1:], }) diff --git a/tool/teleport/testenv/test_server.go b/tool/teleport/testenv/test_server.go index a48cbf16c97..e2a24447e30 100644 --- a/tool/teleport/testenv/test_server.go +++ b/tool/teleport/testenv/test_server.go @@ -27,7 +27,6 @@ import ( "net" "net/http" "net/http/httptest" - "os" "path/filepath" "time" @@ -53,8 +52,6 @@ import ( "github.com/gravitational/teleport/lib/tlsca" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/utils/log/logtest" - "github.com/gravitational/teleport/session/reexec" - "github.com/gravitational/teleport/tool/teleport/common" ) const ( @@ -66,13 +63,6 @@ const ( const StaticToken = "test-static-token" func init() { - // If the test is re-executing itself, execute the command that comes over - // the pipe. Used to test tsh ssh and tsh scp commands. - if reexec.IsReexec() { - common.Run(common.Options{Args: os.Args[1:]}) - return - } - modules.SetModules(&cliModules{}) } diff --git a/tool/tsh/common/tsh_test.go b/tool/tsh/common/tsh_test.go index 36f01e709a2..86d4727150a 100644 --- a/tool/tsh/common/tsh_test.go +++ b/tool/tsh/common/tsh_test.go @@ -125,6 +125,7 @@ var ports utils.PortList const initTestSentinel = "init_test" func TestMain(m *testing.M) { + reexec.MaybeReexec() handleReexec() var err error @@ -238,11 +239,6 @@ func handleReexec() { } os.Exit(0) } - - // Re-exec teleport commands. Used to test tsh ssh command. - if reexec.IsReexec() { - reexec.RunAndExit(os.Args[1]) - } } type cliModules struct{}