diff --git a/go.mod b/go.mod index 097d4443..76f61de3 100644 --- a/go.mod +++ b/go.mod @@ -13,7 +13,7 @@ require ( github.com/pkg/errors v0.9.1 github.com/pkg/sftp v1.13.4 github.com/stretchr/testify v1.7.0 - golang.org/x/crypto v0.0.0-20220307211146-efcb8507fb70 + golang.org/x/crypto v0.40.0 ) require ( @@ -33,6 +33,6 @@ require ( github.com/aws/smithy-go v1.22.4 // indirect github.com/davecgh/go-spew v1.1.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect - golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1 // indirect + golang.org/x/sys v0.34.0 // indirect gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c // indirect ) diff --git a/go.sum b/go.sum index da3fbfcb..0fc4037b 100644 --- a/go.sum +++ b/go.sum @@ -55,8 +55,8 @@ github.com/yuin/goldmark v1.3.5/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1 golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20210421170649-83a5a9bb288b/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4= -golang.org/x/crypto v0.0.0-20220307211146-efcb8507fb70 h1:syTAU9FwmvzEoIYMqcPHOcVm4H3U5u90WsvuYgwpETU= -golang.org/x/crypto v0.0.0-20220307211146-efcb8507fb70/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4= +golang.org/x/crypto v0.40.0 h1:r4x+VvoG5Fm+eJcxMaY8CQM7Lb0l1lsmjGBQ6s8BfKM= +golang.org/x/crypto v0.40.0/go.mod h1:Qr1vMER5WyS2dfPHAlsOj01wgLbsyWtFn/aY+5+ZdxY= golang.org/x/mod v0.4.2/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= @@ -70,10 +70,11 @@ golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.0.0-20210330210617-4fbd30eecc44/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210423185535-09eb48e85fd7/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210510120138-977fb7262007/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1 h1:SrN+KX8Art/Sf4HNj6Zcz06G7VEz+7w9tdXTPOZ7+l4= -golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1 h1:v+OssWQX+hTHEmOBgwxdZxK4zHq3yOs8F9J7mk0PY8E= +golang.org/x/sys v0.34.0 h1:H5Y5sJ2L2JRdyv7ROF1he/lPdvFsd0mJHFw2ThKHxLA= +golang.org/x/sys v0.34.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= +golang.org/x/term v0.33.0 h1:NuFncQrRcaRvVmgRkvM3j/F00gWIAlcmlB8ACEKmGIg= +golang.org/x/term v0.33.0/go.mod h1:s18+ql9tYWp1IfpV9DmCtQDDSRBUjKaw9M1eAv5UeF0= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= diff --git a/managedserver.go b/managedserver.go index 07d90917..02d40717 100644 --- a/managedserver.go +++ b/managedserver.go @@ -1,6 +1,7 @@ package sftp import ( + "errors" "fmt" "io" "net" @@ -28,6 +29,20 @@ type Alerter func(title string, metadata map[string]interface{}) // is valid. type DriverGenerator func(LoginRequest) ServerDriver +// DriverGeneratorWithError is a DriverGenerator that can explain a denial: returning a +// nil driver with a *RejectionError sends the error's message to the SSH client. +type DriverGeneratorWithError func(LoginRequest) (ServerDriver, error) + +// RejectionError rejects a login with a message shown to the SSH client as an +// authentication banner (most clients print it before "Permission denied"). +type RejectionError struct { + Message string +} + +func (e *RejectionError) Error() string { + return e.Message +} + // LoginRequest is the metadata associated with a login request that is passed to the // driverGenerator function in order for it to approve/deny the request. type LoginRequest struct { @@ -43,7 +58,7 @@ type subsystemRequest struct { // ManagedServer is our term for the SFTP server. type ManagedServer struct { - driverGenerator func(LoginRequest) ServerDriver + driverGenerator DriverGeneratorWithError lg Logger alertFn Alerter } @@ -51,6 +66,14 @@ type ManagedServer struct { // NewManagedServer creates a new ManagedServer which conditionally serves requests based // on the output of driverGenerator. func NewManagedServer(driverGenerator DriverGenerator, lg Logger, alertFn Alerter) *ManagedServer { + return NewManagedServerWithError(func(r LoginRequest) (ServerDriver, error) { + return driverGenerator(r), nil + }, lg, alertFn) +} + +// NewManagedServerWithError is NewManagedServer for a DriverGeneratorWithError, allowing +// denials to carry a client-visible message via *RejectionError. +func NewManagedServerWithError(driverGenerator DriverGeneratorWithError, lg Logger, alertFn Alerter) *ManagedServer { return &ManagedServer{ driverGenerator: driverGenerator, lg: lg, @@ -58,6 +81,20 @@ func NewManagedServer(driverGenerator DriverGenerator, lg Logger, alertFn Alerte } } +// rejectionToSSHError converts a driver-generator denial into the error returned to the +// ssh auth callback, sending a banner to the client when a *RejectionError is provided. +func rejectionToSSHError(genErr error, username string) error { + var rejection *RejectionError + if errors.As(genErr, &rejection) { + message := rejection.Message + if !strings.HasSuffix(message, "\n") { + message += "\n" + } + return &ssh.BannerError{Err: genErr, Message: message} + } + return fmt.Errorf("password rejected for %q", username) +} + func (m ManagedServer) errorAndAlert(title string, metadata map[string]interface{}) { if m.alertFn != nil { m.alertFn(title, metadata) @@ -109,26 +146,28 @@ func (m ManagedServer) Start(port int, rawPrivateKeys [][]byte, ciphers, macs [] MACs: macs, }, PasswordCallback: func(c ssh.ConnMetadata, pass []byte) (*ssh.Permissions, error) { - driver = m.driverGenerator(LoginRequest{ + var genErr error + driver, genErr = m.driverGenerator(LoginRequest{ Username: c.User(), Password: string(pass), PublicKey: "", RemoteAddr: c.RemoteAddr(), }) if driver == nil { - return nil, fmt.Errorf("password rejected for %q", c.User()) + return nil, rejectionToSSHError(genErr, c.User()) } return nil, nil }, PublicKeyCallback: func(c ssh.ConnMetadata, key ssh.PublicKey) (*ssh.Permissions, error) { - driver = m.driverGenerator(LoginRequest{ + var genErr error + driver, genErr = m.driverGenerator(LoginRequest{ Username: c.User(), Password: "", PublicKey: strings.TrimSpace(string(ssh.MarshalAuthorizedKey(key))), RemoteAddr: c.RemoteAddr(), }) if driver == nil { - return nil, fmt.Errorf("password rejected for %q", c.User()) + return nil, rejectionToSSHError(genErr, c.User()) } return nil, nil }, diff --git a/managedserver_test.go b/managedserver_test.go new file mode 100644 index 00000000..1354aee2 --- /dev/null +++ b/managedserver_test.go @@ -0,0 +1,58 @@ +package sftp + +import ( + "fmt" + "testing" + + "github.com/stretchr/testify/assert" + "golang.org/x/crypto/ssh" +) + +func TestRejectionToSSHError(t *testing.T) { + for _, spec := range []struct { + desc string + genErr error + wantBanner string + wantGenericMsg bool + }{ + { + desc: "rejection error becomes banner with trailing newline", + genErr: &RejectionError{Message: "use sftp2.clever.com"}, + wantBanner: "use sftp2.clever.com\n", + }, + { + desc: "rejection error keeps existing trailing newline", + genErr: &RejectionError{Message: "use sftp2.clever.com\n"}, + wantBanner: "use sftp2.clever.com\n", + }, + { + desc: "wrapped rejection error becomes banner", + genErr: fmt.Errorf("login denied: %w", &RejectionError{Message: "use sftp2.clever.com"}), + wantBanner: "use sftp2.clever.com\n", + }, + { + desc: "nil error yields generic rejection", + genErr: nil, + wantGenericMsg: true, + }, + { + desc: "plain error yields generic rejection", + genErr: fmt.Errorf("some internal failure"), + wantGenericMsg: true, + }, + } { + t.Run(spec.desc, func(t *testing.T) { + err := rejectionToSSHError(spec.genErr, "some-user") + assert.Error(t, err) + + if spec.wantGenericMsg { + assert.Equal(t, `password rejected for "some-user"`, err.Error()) + return + } + + bannerErr, ok := err.(*ssh.BannerError) + assert.True(t, ok, "expected *ssh.BannerError, got %T", err) + assert.Equal(t, spec.wantBanner, bannerErr.Message) + }) + } +}