407c4ec72d1006cee1ff8c1775e5bcc091c2bc89

Author
Ayman Bagabas <ayman.bagabas@gmail.com>
Committer
Ayman Bagabas <ayman.bagabas@gmail.com>
Date

Message

fix(ssh): add authentication middleware

We need to verify that the key used to establish the connection is the
same key used for authentication, otherwise, refuse connection.

Diff

  1diff --git a/server/ssh/middleware.go b/server/ssh/middleware.go
  2index 300dd3798b2b71670dd2e036b0908b400b02a70f..23bf5ea972762f75123f43e1a13a9ebcd27254e3 100644
  3--- a/server/ssh/middleware.go
  4+++ b/server/ssh/middleware.go
  5@@ -13,11 +13,45 @@ import (
  6 	"github.com/charmbracelet/soft-serve/server/sshutils"
  7 	"github.com/charmbracelet/soft-serve/server/store"
  8 	"github.com/charmbracelet/ssh"
  9+	"github.com/charmbracelet/wish"
 10 	"github.com/prometheus/client_golang/prometheus"
 11 	"github.com/prometheus/client_golang/prometheus/promauto"
 12 	"github.com/spf13/cobra"
 13+	gossh "golang.org/x/crypto/ssh"
 14 )
 15 
 16+// ErrPermissionDenied is returned when a user is not allowed connect.
 17+var ErrPermissionDenied = fmt.Errorf("permission denied")
 18+
 19+// AuthenticationMiddleware handles authentication.
 20+func AuthenticationMiddleware(sh ssh.Handler) ssh.Handler {
 21+	return func(s ssh.Session) {
 22+		// XXX: The authentication key is set in the context but gossh doesn't
 23+		// validate the authentication. We need to verify that the _last_ key
 24+		// that was approved is the one that's being used.
 25+
 26+		pk := s.PublicKey()
 27+		if pk != nil {
 28+			// There is no public key stored in the context, public-key auth
 29+			// was never requested, skip
 30+			perms := s.Permissions().Permissions
 31+			if perms == nil {
 32+				wish.Fatalln(s, ErrPermissionDenied)
 33+				return
 34+			}
 35+
 36+			// Check if the key is the same as the one we have in context
 37+			fp := perms.Extensions["pubkey-fp"]
 38+			if fp != gossh.FingerprintSHA256(pk) {
 39+				wish.Fatalln(s, ErrPermissionDenied)
 40+				return
 41+			}
 42+		}
 43+
 44+		sh(s)
 45+	}
 46+}
 47+
 48 // ContextMiddleware adds the config, backend, and logger to the session context.
 49 func ContextMiddleware(cfg *config.Config, dbx *db.DB, datastore store.Store, be *backend.Backend, logger *log.Logger) func(ssh.Handler) ssh.Handler {
 50 	return func(sh ssh.Handler) ssh.Handler {
 51diff --git a/server/ssh/ssh.go b/server/ssh/ssh.go
 52index d02824da59c88a2383639e037db278f9c78f114b..4d57dc90b212c6525b284b94df9bff54bd7d555a 100644
 53--- a/server/ssh/ssh.go
 54+++ b/server/ssh/ssh.go
 55@@ -77,6 +77,11 @@ func NewSSHServer(ctx context.Context) (*SSHServer, error) {
 56 			LoggingMiddleware,
 57 			// Context middleware.
 58 			ContextMiddleware(cfg, dbx, datastore, be, logger),
 59+			// Authentication middleware.
 60+			// gossh.PublicKeyHandler doesn't guarantee that the public key
 61+			// is in fact the one used for authentication, so we need to
 62+			// check it again here.
 63+			AuthenticationMiddleware,
 64 		),
 65 	}
 66 
 67@@ -91,6 +96,16 @@ func NewSSHServer(ctx context.Context) (*SSHServer, error) {
 68 		return nil, err
 69 	}
 70 
 71+	if config.IsDebug() {
 72+		s.srv.ServerConfigCallback = func(ctx ssh.Context) *gossh.ServerConfig {
 73+			return &gossh.ServerConfig{
 74+				AuthLogCallback: func(conn gossh.ConnMetadata, method string, err error) {
 75+					logger.Debug("authentication", "user", conn.User(), "method", method, "err", err)
 76+				},
 77+			}
 78+		}
 79+	}
 80+
 81 	if cfg.SSH.MaxTimeout > 0 {
 82 		s.srv.MaxTimeout = time.Duration(cfg.SSH.MaxTimeout) * time.Second
 83 	}
 84@@ -130,6 +145,19 @@ func (s *SSHServer) Shutdown(ctx context.Context) error {
 85 	return s.srv.Shutdown(ctx)
 86 }
 87 
 88+func initializePermissions(ctx ssh.Context) {
 89+	perms := ctx.Permissions()
 90+	if perms == nil || perms.Permissions == nil {
 91+		perms = &ssh.Permissions{Permissions: &gossh.Permissions{}}
 92+	}
 93+	if perms.Extensions == nil {
 94+		perms.Extensions = make(map[string]string)
 95+	}
 96+	if perms.Permissions.Extensions == nil {
 97+		perms.Permissions.Extensions = make(map[string]string)
 98+	}
 99+}
100+
101 // PublicKeyAuthHandler handles public key authentication.
102 func (s *SSHServer) PublicKeyHandler(ctx ssh.Context, pk ssh.PublicKey) (allowed bool) {
103 	if pk == nil {
104@@ -144,6 +172,15 @@ func (s *SSHServer) PublicKeyHandler(ctx ssh.Context, pk ssh.PublicKey) (allowed
105 	if user != nil {
106 		ctx.SetValue(proto.ContextKeyUser, user)
107 		allowed = true
108+
109+		// XXX: store the first "approved" public-key fingerprint in the
110+		// permissions block to use for authentication later.
111+		initializePermissions(ctx)
112+		perms := ctx.Permissions()
113+
114+		// Set the public key fingerprint to be used for authentication.
115+		perms.Extensions["pubkey-fp"] = gossh.FingerprintSHA256(pk)
116+		ctx.SetValue(ssh.ContextKeyPermissions, perms)
117 	}
118 
119 	return
120@@ -154,5 +191,16 @@ func (s *SSHServer) PublicKeyHandler(ctx ssh.Context, pk ssh.PublicKey) (allowed
121 func (s *SSHServer) KeyboardInteractiveHandler(ctx ssh.Context, _ gossh.KeyboardInteractiveChallenge) bool {
122 	ac := s.be.AllowKeyless(ctx)
123 	keyboardInteractiveCounter.WithLabelValues(strconv.FormatBool(ac)).Inc()
124+
125+	// If we're allowing keyless access, reset the public key fingerprint
126+	if ac {
127+		initializePermissions(ctx)
128+		perms := ctx.Permissions()
129+
130+		// XXX: reset the public-key fingerprint. This is used to validate the
131+		// public key being used to authenticate.
132+		perms.Extensions["pubkey-fp"] = ""
133+		ctx.SetValue(ssh.ContextKeyPermissions, perms)
134+	}
135 	return ac
136 }