From ce4f41504525e697d0dc5cfbce12ccd83adac328 Mon Sep 17 00:00:00 2001 From: rosstimothy <39066650+rosstimothy@users.noreply.github.com> Date: Wed, 3 Jan 2024 16:37:20 -0500 Subject: [PATCH] Refactor hostname resolution for SSH connections via the WebUI (#35773) * Refactor hostname resolution for SSH connections via the WebUI Previously all hostname resolution was performed by sending either the hostname or server uuid to Auth in a GetSSHTargets request. The returned information was then sent to the UI so that the console tab can be updated to include `user@host`. While this worked, it required a round trip to Auth that had to iterate through the entire set of nodes to do the resolution. During a concurrent session test that tried to spawn several thousand web sessions it was discovered that this cause a massive spike in CPU on Auth. There are two reasons that Auth was used to resolve the hostname. 1) The proxy doesn't perform RBAC on behalf of the user. Instead the proxy has an Auth client with the users identity. 2) The UI used to provide an input box that users could manually enter the target host. Since we no longer allow open dialing and the input box to enter a connection string manually has been removed we no longer need to worry about the second case. To avoid the round trip to auth we can use the local Proxy cache to look up the hostname if we wait until AFTER the ssh connection to the target is established. At this point we can be sure that the user does have access to the target since the node allowed the connection. * Fix terminal tests The testing utilities used to mimic the web ui half of the websocket made assumptions that the session metadata would be the first message provided. However, now that hostname resolution occurs after the initial SSH connection has been established it is possible for other messages to be sent first. This caused some tests to fail and others to hang forever. The duplicated terminal logic from (WebSuite) makeTerminal and (testProxy makeTerminal has now consolidated into a single test terminal utility. The new `terminal` now wraps a TerminalStream which allows test code to make use of the message processing logic already in place there instead of having to write custom logic per test. All a test needs to do is provide handlers for any messages that it desires to introspect. --- integration/helpers/instance.go | 30 +- lib/benchmark/web.go | 3 +- lib/web/apiserver.go | 45 +- lib/web/apiserver_test.go | 789 +++++++++++--------------------- lib/web/command.go | 29 +- lib/web/terminal.go | 162 ++++--- lib/web/terminal_test.go | 180 +++++++- 7 files changed, 590 insertions(+), 648 deletions(-) diff --git a/integration/helpers/instance.go b/integration/helpers/instance.go index de6fcef6c33..4b5db259978 100644 --- a/integration/helpers/instance.go +++ b/integration/helpers/instance.go @@ -36,7 +36,6 @@ import ( "testing" "time" - "github.com/gogo/protobuf/proto" "github.com/gorilla/websocket" "github.com/gravitational/roundtrip" "github.com/gravitational/trace" @@ -64,7 +63,6 @@ import ( "github.com/gravitational/teleport/lib/service" "github.com/gravitational/teleport/lib/service/servicecfg" "github.com/gravitational/teleport/lib/services" - "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/sshutils" "github.com/gravitational/teleport/lib/tlsca" "github.com/gravitational/teleport/lib/utils" @@ -1552,33 +1550,7 @@ func (w *WebClient) SSH(termReq web.TerminalRequest) (*web.TerminalStream, error } defer resp.Body.Close() - ty, raw, err := ws.ReadMessage() - if err != nil { - return nil, trace.Wrap(err) - } - - if ty != websocket.BinaryMessage { - return nil, trace.BadParameter("unexpected websocket message; got %d want %d", ty, websocket.BinaryMessage) - } - - var env web.Envelope - err = proto.Unmarshal(raw, &env) - if err != nil { - return nil, trace.Wrap(err) - } - - type siteSessionGenerateResponse struct { - Session session.Session `json:"session"` - } - - var sessResp siteSessionGenerateResponse - err = json.Unmarshal([]byte(env.Payload), &sessResp) - if err != nil { - return nil, trace.Wrap(err) - } - - stream := web.NewTerminalStream(context.Background(), ws, utils.NewLoggerForTests()) - return stream, nil + return web.NewTerminalStream(context.Background(), web.TerminalStreamConfig{WS: ws}), nil } // AddClientCredentials adds authenticated credentials to a client. diff --git a/lib/benchmark/web.go b/lib/benchmark/web.go index 3e13199d541..7e6dce8cae5 100644 --- a/lib/benchmark/web.go +++ b/lib/benchmark/web.go @@ -39,7 +39,6 @@ import ( "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/client" "github.com/gravitational/teleport/lib/session" - "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/web" ) @@ -193,7 +192,7 @@ func connectToHost(ctx context.Context, tc *client.TeleportClient, webSession *w return nil, trace.BadParameter("unexpected websocket message received %d", ty) } - stream := web.NewTerminalStream(ctx, ws, utils.NewLogger()) + stream := web.NewTerminalStream(ctx, web.TerminalStreamConfig{WS: ws}) return stream, trace.Wrap(err) } diff --git a/lib/web/apiserver.go b/lib/web/apiserver.go index 1e9b0677db2..93b11fc348b 100644 --- a/lib/web/apiserver.go +++ b/lib/web/apiserver.go @@ -2932,7 +2932,7 @@ func (h *Handler) siteNodeConnect( clusterName := site.GetName() if req.SessionID.IsZero() { // An existing session ID was not provided so we need to create a new one. - sessionData, err = h.generateSession(ctx, clt, &req, clusterName, sessionCtx) + sessionData, err = h.generateSession(&req, clusterName, sessionCtx) if err != nil { h.log.WithError(err).Debug("Unable to generate new ssh session.") return nil, trace.Wrap(err) @@ -2980,8 +2980,8 @@ func (h *Handler) siteNodeConnect( terminalConfig := TerminalHandlerConfig{ Term: req.Term, SessionCtx: sessionCtx, - AuthProvider: clt, - LocalAuthProvider: h.auth.accessPoint, + UserAuthClient: clt, + LocalAccessPoint: h.auth.accessPoint, DisplayLogin: displayLogin, SessionData: sessionData, KeepAliveInterval: keepAliveInterval, @@ -3012,11 +3012,11 @@ func (h *Handler) siteNodeConnect( return nil, nil } -func (h *Handler) generateSession(ctx context.Context, clt auth.ClientI, req *TerminalRequest, clusterName string, scx *SessionContext) (session.Session, error) { +func (h *Handler) generateSession(req *TerminalRequest, clusterName string, scx *SessionContext) (session.Session, error) { owner := scx.cfg.User h.log.Infof("Generating new session for %s\n", clusterName) - host, err := findByHost(ctx, clt, req.Server) + host, port, err := serverHostPort(req.Server) if err != nil { return session.Session{}, trace.Wrap(err) } @@ -3031,10 +3031,10 @@ func (h *Handler) generateSession(ctx context.Context, clt auth.ClientI, req *Te return session.Session{ Kind: types.SSHSessionKind, Login: req.Login, - ServerID: host.id, + ServerID: host, ClusterName: clusterName, - ServerHostname: host.hostName, - ServerHostPort: host.port, + ServerHostname: host, + ServerHostPort: port, Moderated: accessEvaluator.IsModerated(), ID: session.NewID(), Created: time.Now().UTC(), @@ -3093,35 +3093,6 @@ func findByQuery(ctx context.Context, clt auth.ClientI, query string) ([]hostInf return hosts, nil } -// findByHost return a host matching by the host name. -func findByHost(ctx context.Context, clt auth.ClientI, serverName string) (*hostInfo, error) { - initialHost, initialPort, err := serverHostPort(serverName) - if err != nil { - return nil, trace.Wrap(err) - } - - rsp, err := clt.GetSSHTargets(ctx, &proto.GetSSHTargetsRequest{ - Host: initialHost, - Port: strconv.Itoa(initialPort), - }) - if err != nil { - return nil, trace.Wrap(err) - } - - var host hostInfo - if len(rsp.Servers) == 1 { - host.hostName = rsp.Servers[0].GetHostname() - host.id = rsp.Servers[0].GetName() - host.port = 0 - } else { - host.hostName = initialHost - host.id = initialHost - host.port = initialPort - } - - return &host, nil -} - // fetchExistingSession fetches an active or pending SSH session by the SessionID passed in the TerminalRequest. func (h *Handler) fetchExistingSession(ctx context.Context, clt auth.ClientI, req *TerminalRequest, siteName string) (session.Session, types.SessionTracker, error) { sessionID, err := session.ParseID(req.SessionID.String()) diff --git a/lib/web/apiserver_test.go b/lib/web/apiserver_test.go index e6e735e3ff3..0bcc359bc1b 100644 --- a/lib/web/apiserver_test.go +++ b/lib/web/apiserver_test.go @@ -1289,8 +1289,18 @@ func TestClusterAlertsGet(t *testing.T) { func TestSiteNodeConnectInvalidSessionID(t *testing.T) { t.Parallel() s := newWebSuite(t) - _, _, err := s.makeTerminal(t, s.authPack(t, "foo"), withSessionID("/../../../foo")) + + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) + + term, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "foo"), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + sessionID: "/../../../foo", + }) require.Error(t, err) + require.Nil(t, term) } func TestResolveServerHostPort(t *testing.T) { @@ -1399,43 +1409,23 @@ func TestFileTransferEvents(t *testing.T) { t.Parallel() s := newWebSuiteWithConfig(t, webSuiteConfig{disableDiskBasedRecording: true}) - errs := make(chan error, 2) - readLoop := func(ctx context.Context, ws *websocket.Conn, ch chan<- *Envelope) { - for { - select { - case <-ctx.Done(): - return - default: - } - - typ, b, err := ws.ReadMessage() - if err != nil { - errs <- err - return - } - if typ != websocket.BinaryMessage { - errs <- trace.BadParameter("expected binary message, got %v", typ) - return - } - var envelope Envelope - if err := proto.Unmarshal(b, &envelope); err != nil { - errs <- trace.Wrap(err) - return - } - ch <- &envelope - } - } + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) // Create a new user "foo", open a terminal to a new session - pack := s.authPack(t, "foo") - ws, _, err := s.makeTerminal(t, pack) - require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws.Close()) }) - - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) wsMessages := make(chan *Envelope) - go readLoop(ctx, ws, wsMessages) + term, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "foo"), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + handlers: map[string]WSHandlerFunc{ + defaults.WebsocketAudit: func(ctx context.Context, envelope Envelope) { + wsMessages <- &envelope + }, + }, + }) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, term.Close()) }) // Create file transfer event data, err := json.Marshal(events.EventFields{ @@ -1451,7 +1441,7 @@ func TestFileTransferEvents(t *testing.T) { } envelopeBytes, err := proto.Marshal(envelope) require.NoError(t, err) - err = ws.WriteMessage(websocket.BinaryMessage, envelopeBytes) + err = term.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes) require.NoError(t, err) done := time.After(5 * time.Second) @@ -1459,8 +1449,6 @@ func TestFileTransferEvents(t *testing.T) { select { case <-done: require.FailNow(t, "expected to receive a file transfer event") - case err := <-errs: - require.NoError(t, err) case e := <-wsMessages: if isFileTransferRequest(e) { requestId, err := getRequestId(e) @@ -1477,7 +1465,7 @@ func TestFileTransferEvents(t *testing.T) { } envelopeBytes, err := proto.Marshal(envelope) require.NoError(t, err) - err = ws.WriteMessage(websocket.BinaryMessage, envelopeBytes) + err = term.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes) require.NoError(t, err) } @@ -1568,10 +1556,10 @@ func TestNewTerminalHandler(t *testing.T) { H: 100, }, SessionCtx: &SessionContext{}, - AuthProvider: authProviderMock{ + UserAuthClient: authProviderMock{ server: validNode, }, - LocalAuthProvider: authProviderMock{}, + LocalAccessPoint: authProviderMock{}, SessionData: session.Session{ ID: session.NewID(), Login: "root", @@ -1588,7 +1576,7 @@ func TestNewTerminalHandler(t *testing.T) { require.NoError(t, err) // passed through require.Equal(t, validCfg.SessionCtx, term.ctx) - require.Equal(t, validCfg.AuthProvider, term.authProvider) + require.Equal(t, validCfg.UserAuthClient, term.userAuthClient) require.Equal(t, validCfg.SessionData, term.sessionData) require.Equal(t, validCfg.KeepAliveInterval, term.keepAliveInterval) require.Equal(t, validCfg.ProxyHostPort, term.proxyHostPort) @@ -1628,100 +1616,95 @@ func TestResizeTerminal(t *testing.T) { s := newWebSuiteWithConfig(t, webSuiteConfig{disableDiskBasedRecording: true}) sid := session.NewID() - errs := make(chan error, 2) - readLoop := func(ctx context.Context, ws *websocket.Conn, ch chan<- *Envelope) { - for { - select { - case <-ctx.Done(): - return - default: - } - - typ, b, err := ws.ReadMessage() - if err != nil { - errs <- err - return - } - if typ != websocket.BinaryMessage { - errs <- trace.BadParameter("expected binary message, got %v", typ) - return - } - var envelope Envelope - if err := proto.Unmarshal(b, &envelope); err != nil { - errs <- trace.Wrap(err) - return - } - ch <- &envelope - } - } + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) + ws1Messages := make(chan *Envelope) + ws1Raw := make(chan []byte) + ws2Messages := make(chan *Envelope) // Create a new user "foo", open a terminal to a new session - pack1 := s.authPack(t, "foo") - ws1, sess, err := s.makeTerminal(t, pack1) + term, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "foo"), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + handlers: map[string]WSHandlerFunc{ + defaults.WebsocketAudit: func(ctx context.Context, envelope Envelope) { + ws1Messages <- &envelope + }, + }, + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws1.Close()) }) + t.Cleanup(func() { require.NoError(t, term.Close()) }) + sess := term.GetSession() // Wait for session to have started require.Eventually(t, func() bool { - _, err := s.server.Auth().GetSessionTracker(context.Background(), sess.ID.String()) + _, err := s.server.Auth().GetSessionTracker(context.Background(), string(sess.ID)) return err == nil }, 3*time.Second, 200*time.Millisecond, "session not available") // Create a new user "bar" and join the session created above - pack2 := s.authPack(t, "bar") - ws2, sess2, err := s.makeTerminal(t, pack2, withSessionID(sess.ID), withParticipantMode(types.SessionPeerMode)) + term2, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "bar"), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + sessionID: sess.ID, + participantMode: types.SessionPeerMode, + handlers: map[string]WSHandlerFunc{ + defaults.WebsocketAudit: func(ctx context.Context, envelope Envelope) { + ws2Messages <- &envelope + }, + }, + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws2.Close()) }) + t.Cleanup(func() { require.NoError(t, term2.Close()) }) - require.Equal(t, sess.ID, sess2.ID) + require.Equal(t, sess.ID, term2.GetSession().ID) - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) - ws1Messages := make(chan *Envelope) - ws2Messages := make(chan *Envelope) - go readLoop(ctx, ws1, ws1Messages) - go readLoop(ctx, ws2, ws2Messages) + go func() { + read, err := io.ReadAll(io.LimitReader(term, 10)) + if err != nil { + return + } - // consume events from the first terminal - // we exect to see at least one raw event with PTY data (indicating terminal ready) - // and 2 resize events from the second user joining the session (one for the default - // size, and one for the manual resize request) + ws1Raw <- read + }() + + // Consume events from the first terminal. We expect to see 2 resize events from the second user + // joining the session (one for the default size, and one for the manual resize request). We also + // validate at least one raw event with PTY data (indicating terminal ready) came through. done := time.After(10 * time.Second) t1ResizeEvents, t1RawEvents := 0, 0 t1ready: for { select { case <-done: - require.FailNowf(t, "", "expected to receive 2 resize events (got %d) and at least 1 raw event (got %d)", t1ResizeEvents, t1RawEvents) - case err := <-errs: - require.NoError(t, err) + require.FailNowf(t, "", "expected to receive 2 resize events (got %d)", t1ResizeEvents) + case <-ws1Raw: + t1RawEvents++ case e := <-ws1Messages: if isResizeEventEnvelope(e) { t1ResizeEvents++ } - if e.GetType() == defaults.WebsocketRaw { - t1RawEvents++ - } - if t1ResizeEvents == 2 && t1RawEvents > 0 { - break t1ready - } + } + + if t1ResizeEvents == 2 && t1RawEvents > 0 { + break t1ready } } // we should not expect to see a resize event on terminal 2, - // since they are not broadcasted back to the originator + // since they are not broadcast back to the originator select { case e := <-ws2Messages: if isResizeEventEnvelope(e) { require.FailNow(t, "terminal 2 should not have received a resize event: %v", e) } - case err := <-errs: - require.NoError(t, err) case <-time.After(1 * time.Second): } - // Resize the second terminal. This should be reflected only on the first terminal - // because resize events are sent to participants but not the originator.. + // Resize the second terminal. This should only be reflected in the first terminal + // because resize events are sent to participants but not the originator. params, err := session.NewTerminalParamsFromInt(300, 120) require.NoError(t, err) data, err := json.Marshal(events.EventFields{ @@ -1738,7 +1721,7 @@ t1ready: } envelopeBytes, err := proto.Marshal(envelope) require.NoError(t, err) - err = ws2.WriteMessage(websocket.BinaryMessage, envelopeBytes) + err = term2.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes) require.NoError(t, err) // the first terminal should see the resize event @@ -1747,8 +1730,6 @@ t1ready: select { case <-done: require.FailNow(t, "expected to receive a final resize event") - case err := <-errs: - require.NoError(t, err) case e := <-ws1Messages: if isResizeEventEnvelope(e) { return @@ -1772,37 +1753,38 @@ func isResizeEventEnvelope(e *Envelope) bool { func TestTerminalPing(t *testing.T) { t.Parallel() s := newWebSuite(t) - ws, _, err := s.makeTerminal(t, s.authPack(t, "foo"), withKeepaliveInterval(time.Second)) + + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) + + term, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "foo"), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + keepAliveInterval: time.Second, + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws.Close()) }) + t.Cleanup(func() { require.NoError(t, term.Close()) }) closed := false done := make(chan struct{}) - ws.SetPingHandler(func(message string) error { + term.ws.SetPingHandler(func(message string) error { if closed == false { close(done) closed = true } - err := ws.WriteControl(websocket.PongMessage, []byte(message), time.Now().Add(time.Second)) - if err == websocket.ErrCloseSent { + err := term.ws.WriteControl(websocket.PongMessage, []byte(message), time.Now().Add(time.Second)) + if errors.Is(err, websocket.ErrCloseSent) { return nil - } else if e, ok := err.(net.Error); ok && e.Timeout() { - return nil - } - return err - }) - - // We need to continuously read incoming messages in order to process ping messages. - // We only care about receiving a ping here so dropping them is fine. - go func() { - for { - _, _, err := ws.ReadMessage() - if err != nil { - return + } else { + var e net.Error + if errors.As(err, &e) && e.Timeout() { + return nil } + return err } - }() + }) select { case <-done: @@ -1845,19 +1827,25 @@ func TestTerminal(t *testing.T) { // Set the recording config require.NoError(t, s.server.Auth().SetSessionRecordingConfig(context.Background(), &tt.recordingConfig)) + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) // Create a new session - ws, _, err := s.makeTerminal(t, s.authPack(t, "foo")) + term, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "foo"), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + }) require.NoError(t, err) - t.Cleanup(func() { require.True(t, utils.IsOKNetworkError(ws.Close())) }) + t.Cleanup(func() { require.True(t, utils.IsOKNetworkError(term.Close())) }) // Send a command and validate the output - validateTerminalStream(t, ws) + validateTerminal(t, term) // Validate that the session is active on the node require.Equal(t, int32(1), s.node.ActiveConnections()) // Close the web socket to emulate a user closing the browser window - require.NoError(t, ws.Close()) + require.NoError(t, term.Close()) // Validate that the node terminates the session require.EventuallyWithT(t, func(t *assert.CollectT) { @@ -1871,7 +1859,7 @@ func TestTerminalRouting(t *testing.T) { t.Parallel() s := newWebSuite(t) - // add nodes with various conflicting values + // add nodes with conflicting hostnames llama := s.addNode(t, uuid.NewString(), "llama", "127.0.0.1:0") s.addNode(t, uuid.NewString(), "llamas", "127.0.0.1:0") alpaca1 := s.addNode(t, uuid.NewString(), "alpaca", "127.0.0.1:0") @@ -1881,58 +1869,24 @@ func TestTerminalRouting(t *testing.T) { require.NoError(t, err) } - closeOkNetworkError := func(t *testing.T, err error) { - if err == nil { - return - } - - require.True(t, utils.IsOKNetworkError(err), "websocket closure should have return an error indicating that the server already terminated the connection") - } - cases := []struct { name string - target string + target *regular.Server output string wsCloseAssertion func(t *testing.T, err error) }{ { name: "exact match by uuid", - target: llama.ID(), + target: llama, output: "teleport", wsCloseAssertion: closeNoError, }, - { - name: "exact match by hostname", - target: "llama", - output: "teleport", - wsCloseAssertion: closeNoError, - }, - { - name: "exact match by ip", - target: llama.Addr(), - output: "teleport", - wsCloseAssertion: closeNoError, - }, - { - name: "ambiguous host", - target: "alpaca", - output: "error: ambiguous host could match multiple nodes", - // failed resolution results in the server closing the socket first, so expect an ok close error - wsCloseAssertion: closeOkNetworkError, - }, { name: "connect by uuid successful when multiple hostnames match", - target: alpaca1.ID(), + target: alpaca1, output: "teleport", wsCloseAssertion: closeNoError, }, - { - name: "ambiguous ip", - target: "127.0.0.1", - output: "error: ambiguous host could match multiple nodes", - // failed resolution results in the server closing the socket first, so expect an ok close error - wsCloseAssertion: closeOkNetworkError, - }, } for i, tt := range cases { @@ -1940,85 +1894,28 @@ func TestTerminalRouting(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - ws, _, err := s.makeTerminal(t, s.authPack(t, fmt.Sprintf("foo-%d", i)), withServer(tt.target)) - require.NoError(t, err) - t.Cleanup(func() { tt.wsCloseAssertion(t, ws.Close()) }) + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) - stream := NewTerminalStream(s.ctx, ws, utils.NewLoggerForTests()) + term, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, fmt.Sprintf("foo-%d", i)), + host: tt.target.ID(), + proxy: s.webServer.Listener.Addr().String(), + }) + require.NoError(t, err) + t.Cleanup(func() { tt.wsCloseAssertion(t, term.Close()) }) + + sess := term.GetSession() + + metadata := tt.target.TargetMetadata() + require.Equal(t, metadata.ServerID, sess.ServerID) + require.Equal(t, metadata.ServerHostname, sess.ServerHostname) // here we intentionally run a command where the output we're looking // for is not present in the command itself - _, err = io.WriteString(stream, "echo txlxport | sed 's/x/e/g'\r\n") + _, err = io.WriteString(term, "echo txlxport | sed 's/x/e/g'\r\n") require.NoError(t, err) - require.NoError(t, waitForOutput(stream, tt.output)) - }) - } -} - -func TestTerminalNameResolution(t *testing.T) { - t.Parallel() - s := newWebSuite(t) - pack := s.authPack(t, "foo") - - llama := s.addNode(t, uuid.NewString(), "llama", "127.0.0.1:0") - - ctx, cancel := context.WithTimeout(context.Background(), 7*time.Second) - t.Cleanup(cancel) - - // Wait for the node to be registered as the registration is asynchronous. - require.Eventually(t, func() bool { - nodes, err := s.proxyClient.GetNodes(ctx, "default") - assert.NoError(t, err) - - return len(nodes) == 2 // one created by default and llama - }, 5*time.Second, 200*time.Millisecond, "failed to register node") - - tests := []struct { - name string - target string - serverID string - serverHostname string - port int - }{ - { - name: "registered node by name", - target: "llama", - serverID: llama.ID(), - serverHostname: "llama", - }, - { - name: "registered node by address", - target: llama.Addr(), - serverID: llama.ID(), - serverHostname: "llama", - }, - { - name: "direct dial", - target: "root@example.com", - serverID: "root@example.com", - serverHostname: "root@example.com", - }, - { - name: "direct dial with port", - target: "root@example.com:1234", - serverID: "root@example.com", - serverHostname: "root@example.com", - port: 1234, - }, - } - - for _, tt := range tests { - tt := tt - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - - ws, resp, err := s.makeTerminal(t, pack, withServer(tt.target)) - require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws.Close()) }) - - require.Equal(t, tt.serverID, resp.ServerID) - require.Equal(t, tt.serverHostname, resp.ServerHostname) - require.Equal(t, tt.port, resp.ServerHostPort) + require.NoError(t, waitForOutput(term, tt.output)) }) } } @@ -2038,7 +1935,7 @@ func TestTerminalRequireSessionMFA(t *testing.T) { name string getAuthPreference func(t *testing.T) types.AuthPreference registerDevice func(t *testing.T) *auth.TestDevice - getChallengeResponseBytes func(t *testing.T, chal *client.MFAAuthenticateChallenge, testDev *auth.TestDevice) []byte + getChallengeResponseBytes func(t *testing.T, chal client.MFAAuthenticateChallenge, testDev *auth.TestDevice) []byte }{ { name: "with webauthn", @@ -2064,7 +1961,7 @@ func TestTerminalRequireSessionMFA(t *testing.T) { return webauthnDev }, - getChallengeResponseBytes: func(t *testing.T, chal *client.MFAAuthenticateChallenge, testDev *auth.TestDevice) []byte { + getChallengeResponseBytes: func(t *testing.T, chal client.MFAAuthenticateChallenge, testDev *auth.TestDevice) []byte { res, err := testDev.SolveAuthn(&authproto.MFAAuthenticateChallenge{ WebauthnChallenge: wantypes.CredentialAssertionToProto(chal.WebauthnChallenge), }) @@ -2092,28 +1989,24 @@ func TestTerminalRequireSessionMFA(t *testing.T) { dev := tc.registerDevice(t) + termCtx, cancel := context.WithCancel(ctx) + t.Cleanup(cancel) + // Open a terminal to a new session. - ws, _ := proxy.makeTerminal(t, pack, "") - - // Wait for websocket authn challenge event. - ty, raw, err := ws.ReadMessage() - require.NoError(t, err) - require.Equal(t, websocket.BinaryMessage, ty) - var env Envelope - require.NoError(t, proto.Unmarshal(raw, &env)) - - chal := &client.MFAAuthenticateChallenge{} - require.NoError(t, json.Unmarshal([]byte(env.Payload), &chal)) - - // Send response over ws. - stream := NewTerminalStream(ctx, ws, utils.NewLoggerForTests()) - err = stream.ws.WriteMessage(websocket.BinaryMessage, tc.getChallengeResponseBytes(t, chal, dev)) + term, err := connectToHost(termCtx, connectConfig{ + pack: pack, + host: proxy.node.ID(), + proxy: proxy.webURL.Host, + mfaCeremony: func(challenge client.MFAAuthenticateChallenge) []byte { + return tc.getChallengeResponseBytes(t, challenge, dev) + }, + }) require.NoError(t, err) // Test we can write. - _, err = io.WriteString(stream, "echo txlxport | sed 's/x/e/g'\r\n") + _, err = io.WriteString(term, "echo txlxport | sed 's/x/e/g'\r\n") require.NoError(t, err) - require.NoError(t, waitForOutput(stream, "teleport")) + require.NoError(t, waitForOutput(term, "teleport")) }) } } @@ -2299,16 +2192,22 @@ func handleDesktopMFAWebauthnChallenge(t *testing.T, ws *websocket.Conn, dev *au func TestWebAgentForward(t *testing.T) { t.Parallel() s := newWebSuiteWithConfig(t, webSuiteConfig{disableDiskBasedRecording: true}) - ws, _, err := s.makeTerminal(t, s.authPack(t, "foo")) + + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) + + term, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "foo"), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws.Close()) }) + t.Cleanup(func() { require.NoError(t, term.Close()) }) - stream := NewTerminalStream(s.ctx, ws, utils.NewLoggerForTests()) - - _, err = io.WriteString(stream, "echo $SSH_AUTH_SOCK\r\n") + _, err = io.WriteString(term, "echo $SSH_AUTH_SOCK\r\n") require.NoError(t, err) - err = waitForOutput(stream, "/") + err = waitForOutput(term, "/") require.NoError(t, err) } @@ -2411,19 +2310,24 @@ func TestCloseConnectionsOnLogout(t *testing.T) { s := newWebSuite(t) pack := s.authPack(t, "foo") - ws, _, err := s.makeTerminal(t, pack) - require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws.Close()) }) + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) - stream := NewTerminalStream(s.ctx, ws, utils.NewLoggerForTests()) + term, err := connectToHost(ctx, connectConfig{ + pack: pack, + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + }) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, term.Close()) }) // to make sure we have a session - _, err = io.WriteString(stream, "expr 137 + 39\r\n") + _, err = io.WriteString(term, "expr 137 + 39\r\n") require.NoError(t, err) // make sure the server has replied out := make([]byte, 100) - _, err = stream.Read(out) + _, err = term.Read(out) require.NoError(t, err) _, err = pack.clt.Delete(s.ctx, pack.clt.Endpoint("webapi", "sessions", "web")) @@ -2434,7 +2338,7 @@ func TestCloseConnectionsOnLogout(t *testing.T) { errC := make(chan error) go func() { for { - _, err := stream.Read(out) + _, err := term.Read(out) if err != nil { errC <- err return @@ -2450,15 +2354,6 @@ func TestCloseConnectionsOnLogout(t *testing.T) { } } -func TestPlayback(t *testing.T) { - t.Parallel() - s := newWebSuite(t) - pack := s.authPack(t, "foo") - ws, _, err := s.makeTerminal(t, pack) - require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, ws.Close()) }) -} - type httpErrorMessage struct { Message string `json:"message"` } @@ -5311,7 +5206,16 @@ func TestWebSessionsRenewDoesNotBreakExistingTerminalSession(t *testing.T) { pack1 := proxy1.authPack(t, "foo", nil /* roles */) pack2 := proxy2.authPackFromPack(t, pack1) - ws, _ := proxy2.makeTerminal(t, pack2, "") + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + + term, err := connectToHost(ctx, connectConfig{ + pack: pack2, + host: proxy2.node.ID(), + proxy: proxy2.webURL.Host, + }) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, term.Close()) }) // Advance the time before renewing the session. // This will allow the new session to have a more plausible @@ -5332,7 +5236,7 @@ func TestWebSessionsRenewDoesNotBreakExistingTerminalSession(t *testing.T) { pack2.validateAPI(context.Background(), t) // Check whether the terminal session is still active - validateTerminalStream(t, ws) + validateTerminal(t, term) } // TestWebSessionsRenewAllowsOldBearerTokenToLinger validates that the @@ -7325,8 +7229,8 @@ type authProviderMock struct { server types.ServerV2 } -func (mock authProviderMock) GetNodes(ctx context.Context, n string) ([]types.Server, error) { - return []types.Server{&mock.server}, nil +func (mock authProviderMock) GetNode(ctx context.Context, namespace, name string) (types.Server, error) { + return &mock.server, nil } func (mock authProviderMock) GetSessionEvents(n string, s session.ID, c int) ([]events.EventFields, error) { @@ -7365,102 +7269,6 @@ func (mock authProviderMock) GetRole(_ context.Context, _ string) (types.Role, e return nil, nil } -type terminalOpt func(t *TerminalRequest) - -func withSessionID(sid session.ID) terminalOpt { - return func(t *TerminalRequest) { t.SessionID = sid } -} - -func withServer(target string) terminalOpt { - return func(t *TerminalRequest) { t.Server = target } -} - -func withKeepaliveInterval(d time.Duration) terminalOpt { - return func(t *TerminalRequest) { t.KeepAliveInterval = d } -} - -func withParticipantMode(m types.SessionParticipantMode) terminalOpt { - return func(t *TerminalRequest) { t.ParticipantMode = m } -} - -func (s *WebSuite) makeTerminal(t *testing.T, pack *authPack, opts ...terminalOpt) (*websocket.Conn, *session.Session, error) { - req := TerminalRequest{ - Server: s.srvID, - Login: pack.login, - Term: session.TerminalParams{ - W: 100, - H: 100, - }, - } - for _, opt := range opts { - opt(&req) - } - - u := url.URL{ - Host: s.url().Host, - Scheme: client.WSS, - Path: fmt.Sprintf("/v1/webapi/sites/%v/connect", currentSiteShortcut), - } - data, err := json.Marshal(req) - if err != nil { - return nil, nil, err - } - - q := u.Query() - q.Set("params", string(data)) - q.Set(roundtrip.AccessTokenQueryParam, pack.session.Token) - u.RawQuery = q.Encode() - - dialer := websocket.Dialer{} - dialer.TLSClientConfig = &tls.Config{ - InsecureSkipVerify: true, - } - - header := http.Header{} - header.Add("Origin", "http://localhost") - for _, cookie := range pack.cookies { - header.Add("Cookie", cookie.String()) - } - - ws, resp, err := dialer.Dial(u.String(), header) - if err != nil { - var sb strings.Builder - sb.WriteString("websocket dial") - if resp != nil { - fmt.Fprintf(&sb, "; status code %v;", resp.StatusCode) - fmt.Fprintf(&sb, "headers: %v; body: ", resp.Header) - io.Copy(&sb, resp.Body) - } - return nil, nil, trace.Wrap(err, sb.String()) - } - - ty, raw, err := ws.ReadMessage() - if err != nil { - return nil, nil, trace.Wrap(err) - } - require.Equal(t, websocket.BinaryMessage, ty) - var env Envelope - - err = proto.Unmarshal(raw, &env) - if err != nil { - return nil, nil, trace.Wrap(err) - } - - var sessResp siteSessionGenerateResponse - - err = json.Unmarshal([]byte(env.Payload), &sessResp) - if err != nil { - return nil, nil, trace.Wrap(err) - } - - err = resp.Body.Close() - if err != nil { - return nil, nil, trace.Wrap(err) - } - - return ws, &sessResp.Session, nil -} - func waitForOutputWithDuration(r io.Reader, substr string, timeout time.Duration) error { timeoutCh := time.After(timeout) @@ -8232,64 +8040,6 @@ func (r *testProxy) newClient(t *testing.T, opts ...roundtrip.ClientParam) *Test return &TestWebClient{clt, t} } -func (r *testProxy) makeTerminal(t *testing.T, pack *authPack, sessionID session.ID) (*websocket.Conn, session.Session) { - u := url.URL{ - Host: r.webURL.Host, - Scheme: client.WSS, - Path: fmt.Sprintf("/v1/webapi/sites/%v/connect", currentSiteShortcut), - } - - requestData := TerminalRequest{ - Server: r.node.ID(), - Login: pack.login, - Term: session.TerminalParams{ - W: 100, - H: 100, - }, - } - - if sessionID != "" { - requestData.SessionID = sessionID - } - - data, err := json.Marshal(requestData) - require.NoError(t, err) - - q := u.Query() - q.Set("params", string(data)) - q.Set(roundtrip.AccessTokenQueryParam, pack.session.Token) - u.RawQuery = q.Encode() - - dialer := websocket.Dialer{} - dialer.TLSClientConfig = &tls.Config{ - InsecureSkipVerify: true, - } - - header := http.Header{} - header.Add("Origin", "http://localhost") - for _, cookie := range pack.cookies { - header.Add("Cookie", cookie.String()) - } - - ws, resp, err := dialer.Dial(u.String(), header) - require.NoError(t, err) - t.Cleanup(func() { - require.NoError(t, ws.Close()) - require.NoError(t, resp.Body.Close()) - }) - - ty, raw, err := ws.ReadMessage() - require.NoError(t, err) - require.Equal(t, websocket.BinaryMessage, ty) - var env Envelope - require.NoError(t, proto.Unmarshal(raw, &env)) - - var sessResp siteSessionGenerateResponse - require.NoError(t, json.Unmarshal([]byte(env.Payload), &sessResp)) - - return ws, sessResp.Session -} - func (r *testProxy) makeDesktopSession(t *testing.T, pack *authPack, sessionID session.ID, addr net.Addr) *websocket.Conn { u := url.URL{ Host: r.webURL.Host, @@ -8342,15 +8092,14 @@ func login(t *testing.T, clt *TestWebClient, cookieToken, reqToken string, reqDa return resp } -func validateTerminalStream(t *testing.T, ws *websocket.Conn) { +func validateTerminal(t *testing.T, term io.ReadWriter) { t.Helper() - stream := NewTerminalStream(context.Background(), ws, utils.NewLoggerForTests()) // here we intentionally run a command where the output we're looking // for is not present in the command itself - _, err := io.WriteString(stream, "echo txlxport | sed 's/x/e/g'\r\n") + _, err := io.WriteString(term, "echo txlxport | sed 's/x/e/g'\r\n") require.NoError(t, err) - require.NoError(t, waitForOutput(stream, "teleport")) + require.NoError(t, waitForOutput(term, "teleport")) } type mockProxySettings struct { @@ -9381,7 +9130,6 @@ func (m mockedPingTestProxy) Ping(ctx context.Context) (authproto.PingResponse, func TestModeratedSession(t *testing.T) { modules.SetTestModules(t, &modules.TestModules{TestBuildType: modules.BuildEnterprise}) - ctx := context.Background() s := newWebSuiteWithConfig(t, webSuiteConfig{disableDiskBasedRecording: true}) peerRole, err := types.NewRole("moderated", types.RoleSpecV6{ @@ -9417,38 +9165,44 @@ func TestModeratedSession(t *testing.T) { moderatorRole, err = s.server.Auth().UpsertRole(s.ctx, moderatorRole) require.NoError(t, err) - peer := s.authPack(t, "foo", peerRole.GetName()) + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) - peerWS, sess, err := s.makeTerminal(t, peer) + peerTerm, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "foo", peerRole.GetName()), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, peerWS.Close()) }) + t.Cleanup(func() { require.NoError(t, peerTerm.Close()) }) - peerStream := NewTerminalStream(ctx, peerWS, utils.NewLoggerForTests()) + require.NoError(t, waitForOutput(peerTerm, "Teleport > User foo joined the session with participant mode: peer.")) - require.NoError(t, waitForOutput(peerStream, "Teleport > User foo joined the session with participant mode: peer.")) - - moderator := s.authPack(t, "bar", moderatorRole.GetName()) - moderatorWS, _, err := s.makeTerminal(t, moderator, withSessionID(sess.ID), withParticipantMode(types.SessionModeratorMode)) + moderatorTerm, err := connectToHost(ctx, connectConfig{ + pack: s.authPack(t, "bar", moderatorRole.GetName()), + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + sessionID: peerTerm.GetSession().ID, + participantMode: types.SessionModeratorMode, + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, moderatorWS.Close()) }) + t.Cleanup(func() { require.NoError(t, moderatorTerm.Close()) }) - moderatorStream := NewTerminalStream(ctx, moderatorWS, utils.NewLoggerForTests()) - - require.NoError(t, waitForOutput(peerStream, "Teleport > Connecting to node over SSH")) + require.NoError(t, waitForOutput(peerTerm, "Teleport > Connecting to node over SSH")) // here we intentionally run a command where the output we're looking // for is not present in the command itself - _, err = io.WriteString(peerStream, "echo llxmx | sed 's/x/a/g'\r\n") + _, err = io.WriteString(peerTerm, "echo llxmx | sed 's/x/a/g'\r\n") require.NoError(t, err) - require.NoError(t, waitForOutput(peerStream, "llama")) - require.NoError(t, waitForOutput(moderatorStream, "llama")) + require.NoError(t, waitForOutput(peerTerm, "llama")) + require.NoError(t, waitForOutput(moderatorTerm, "llama")) // the moderator terminates the session - _, err = io.WriteString(moderatorStream, "t") + _, err = io.WriteString(moderatorTerm, "t") require.NoError(t, err) - require.NoError(t, waitForOutput(moderatorStream, "Stopping session...")) - require.NoError(t, waitForOutput(peerStream, "Process exited with status 255")) + require.NoError(t, waitForOutput(moderatorTerm, "Stopping session...")) + require.NoError(t, waitForOutput(peerTerm, "Process exited with status 255")) } // TestModeratedSessionWithMFA validates the same behavior as TestModeratedSession while @@ -9457,7 +9211,6 @@ func TestModeratedSession(t *testing.T) { // the session is aborted. func TestModeratedSessionWithMFA(t *testing.T) { modules.SetTestModules(t, &modules.TestModules{TestBuildType: modules.BuildEnterprise}) - ctx := context.Background() const RPID = "localhost" @@ -9511,40 +9264,83 @@ func TestModeratedSessionWithMFA(t *testing.T) { peer := s.authPackWithMFA(t, "foo", peerRole) moderator := s.authPackWithMFA(t, "bar", moderatorRole) - peerWS, sess, err := s.makeTerminal(t, peer) + ctx, cancel := context.WithCancel(s.ctx) + t.Cleanup(cancel) + + peerTerm, err := connectToHost(ctx, connectConfig{ + pack: peer, + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + mfaCeremony: func(challenge client.MFAAuthenticateChallenge) []byte { + res, err := peer.device.SolveAuthn(&authproto.MFAAuthenticateChallenge{ + WebauthnChallenge: wantypes.CredentialAssertionToProto(challenge.WebauthnChallenge), + }) + require.NoError(t, err) + + webauthnResBytes, err := json.Marshal(wantypes.CredentialAssertionResponseFromProto(res.GetWebauthn())) + require.NoError(t, err) + + envelope := &Envelope{ + Version: defaults.WebsocketVersion, + Type: defaults.WebsocketWebauthnChallenge, + Payload: string(webauthnResBytes), + } + envelopeBytes, err := proto.Marshal(envelope) + require.NoError(t, err) + + return envelopeBytes + }, + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, peerWS.Close()) }) + t.Cleanup(func() { require.NoError(t, peerTerm.Close()) }) - handleMFAWebauthnChallenge(t, peerWS, peer.device) + require.NoError(t, waitForOutput(peerTerm, "Teleport > User foo joined the session with participant mode: peer.")) - peerStream := NewTerminalStream(ctx, peerWS, utils.NewLoggerForTests()) + moderatorTerm, err := connectToHost(ctx, connectConfig{ + pack: moderator, + host: s.node.ID(), + proxy: s.webServer.Listener.Addr().String(), + sessionID: peerTerm.GetSession().ID, + participantMode: types.SessionModeratorMode, + mfaCeremony: func(challenge client.MFAAuthenticateChallenge) []byte { + res, err := moderator.device.SolveAuthn(&authproto.MFAAuthenticateChallenge{ + WebauthnChallenge: wantypes.CredentialAssertionToProto(challenge.WebauthnChallenge), + }) + require.NoError(t, err) - require.NoError(t, waitForOutput(peerStream, "Teleport > User foo joined the session with participant mode: peer.")) + webauthnResBytes, err := json.Marshal(wantypes.CredentialAssertionResponseFromProto(res.GetWebauthn())) + require.NoError(t, err) - moderatorWS, _, err := s.makeTerminal(t, moderator, withSessionID(sess.ID), withParticipantMode(types.SessionModeratorMode)) + envelope := &Envelope{ + Version: defaults.WebsocketVersion, + Type: defaults.WebsocketWebauthnChallenge, + Payload: string(webauthnResBytes), + } + envelopeBytes, err := proto.Marshal(envelope) + require.NoError(t, err) + + return envelopeBytes + }, + }) require.NoError(t, err) - t.Cleanup(func() { require.NoError(t, moderatorWS.Close()) }) + t.Cleanup(func() { require.NoError(t, moderatorTerm.Close()) }) - handleMFAWebauthnChallenge(t, moderatorWS, moderator.device) - - moderatorStream := NewTerminalStream(ctx, moderatorWS, utils.NewLoggerForTests()) - - require.NoError(t, waitForOutput(peerStream, "Teleport > Connecting to node over SSH")) + require.NoError(t, waitForOutput(peerTerm, "Teleport > Connecting to node over SSH")) // here we intentionally run a command where the output we're looking // for is not present in the command itself - _, err = io.WriteString(peerStream, "echo llxmx | sed 's/x/a/g'\r\n") + _, err = io.WriteString(peerTerm, "echo llxmx | sed 's/x/a/g'\r\n") require.NoError(t, err) - require.NoError(t, waitForOutput(peerStream, "llama")) - require.NoError(t, waitForOutput(moderatorStream, "llama")) + require.NoError(t, waitForOutput(peerTerm, "llama")) + require.NoError(t, waitForOutput(moderatorTerm, "llama")) // run the presence check a few times for i := 0; i < 3; i++ { presenceClock.BlockUntil(1) presenceClock.Advance(30 * time.Second) - require.NoError(t, waitForOutput(moderatorStream, "Teleport > Please tap your MFA key")) + require.NoError(t, waitForOutput(moderatorTerm, "Teleport > Please tap your MFA key")) - challenge, err := moderatorStream.readChallenge(protobufMFACodec{}) + challenge, err := moderatorTerm.stream.readChallenge(protobufMFACodec{}) require.NoError(t, err) res, err := moderator.device.SolveAuthn(challenge) @@ -9561,7 +9357,7 @@ func TestModeratedSessionWithMFA(t *testing.T) { envelopeBytes, err := proto.Marshal(envelope) require.NoError(t, err) - require.NoError(t, moderatorWS.WriteMessage(websocket.BinaryMessage, envelopeBytes)) + require.NoError(t, moderatorTerm.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes)) } // Advance the clock far enough in the future to make the moderator stale @@ -9569,42 +9365,11 @@ func TestModeratedSessionWithMFA(t *testing.T) { // components, it's not practical to use BlockUntil here, so we use EventuallyWithT instead. require.EventuallyWithT(t, func(t *assert.CollectT) { s.clock.Advance(3 * time.Minute) - assert.NoError(t, waitForOutputWithDuration(moderatorStream, "wait: remote command exited without exit status or exit signal", 3*time.Second)) - assert.NoError(t, waitForOutputWithDuration(peerStream, "Process exited with status 255", 3*time.Second)) + assert.NoError(t, waitForOutputWithDuration(moderatorTerm, "wait: remote command exited without exit status or exit signal", 3*time.Second)) + assert.NoError(t, waitForOutputWithDuration(peerTerm, "Process exited with status 255", 3*time.Second)) }, 15*time.Second, 500*time.Millisecond) } -func handleMFAWebauthnChallenge(t *testing.T, ws *websocket.Conn, dev *auth.TestDevice) { - // Wait for websocket authn challenge event. - ty, raw, err := ws.ReadMessage() - require.NoError(t, err) - require.Equal(t, websocket.BinaryMessage, ty) - - var env Envelope - require.NoError(t, proto.Unmarshal(raw, &env)) - - var challenge client.MFAAuthenticateChallenge - require.NoError(t, json.Unmarshal([]byte(env.Payload), &challenge)) - - res, err := dev.SolveAuthn(&authproto.MFAAuthenticateChallenge{ - WebauthnChallenge: wantypes.CredentialAssertionToProto(challenge.WebauthnChallenge), - }) - require.NoError(t, err) - - webauthnResBytes, err := json.Marshal(wantypes.CredentialAssertionResponseFromProto(res.GetWebauthn())) - require.NoError(t, err) - - envelope := &Envelope{ - Version: defaults.WebsocketVersion, - Type: defaults.WebsocketWebauthnChallenge, - Payload: string(webauthnResBytes), - } - envelopeBytes, err := proto.Marshal(envelope) - require.NoError(t, err) - - require.NoError(t, ws.WriteMessage(websocket.BinaryMessage, envelopeBytes)) -} - type proxyClientMock struct { auth.ClientI tokens map[string]types.ProvisionToken diff --git a/lib/web/command.go b/lib/web/command.go index 7dbd19be8d1..9017b63ed1b 100644 --- a/lib/web/command.go +++ b/lib/web/command.go @@ -239,14 +239,14 @@ func (h *Handler) executeCommand( commandHandlerConfig := CommandHandlerConfig{ SessionCtx: sessionCtx, - AuthProvider: clt, + UserAuthClient: clt, SessionData: sessionData, KeepAliveInterval: keepAliveInterval, ProxyHostPort: h.ProxyHostPort(), InteractiveCommand: interactiveCommand, Router: h.cfg.Router, TracerProvider: h.cfg.TracerProvider, - LocalAuthProvider: h.auth.accessPoint, + LocalAccessPoint: h.auth.accessPoint, mfaFuncCache: mfaCacheFn, buffer: buffer, } @@ -472,13 +472,13 @@ func newCommandHandler(ctx context.Context, cfg CommandHandlerConfig) (*commandH "session_id": cfg.SessionData.ID.String(), }), ctx: cfg.SessionCtx, - authProvider: cfg.AuthProvider, + userAuthClient: cfg.UserAuthClient, sessionData: cfg.SessionData, keepAliveInterval: cfg.KeepAliveInterval, proxyHostPort: cfg.ProxyHostPort, interactiveCommand: cfg.InteractiveCommand, router: cfg.Router, - localAuthProvider: cfg.LocalAuthProvider, + localAccessPoint: cfg.LocalAccessPoint, tracer: cfg.tracer, }, mfaAuthCache: cfg.mfaFuncCache, @@ -490,8 +490,8 @@ func newCommandHandler(ctx context.Context, cfg CommandHandlerConfig) (*commandH type CommandHandlerConfig struct { // SessionCtx is the context for the user's web session. SessionCtx *SessionContext - // AuthProvider is used to fetch nodes and sessions from the backend. - AuthProvider AuthProvider + // UserAuthClient is used to fetch nodes and sessions from the backend via the users' identity. + UserAuthClient UserAuthClient // SessionData is the data to send to the client on the initial session creation. SessionData session.Session // KeepAliveInterval is the interval for sending ping frames to a web client. @@ -506,9 +506,12 @@ type CommandHandlerConfig struct { Router *proxy.Router // TracerProvider is used to create the tracer TracerProvider oteltrace.TracerProvider - // LocalAuthProvider is used to fetch user information from the - // local cluster when connecting to agentless nodes. - LocalAuthProvider agentless.AuthProvider + // LocalAccessPoint is the subset of the Proxy cache required to + // look up information from the local cluster. This should not + // be used for anything that requires RBAC on behalf of the user. + // Anything requests that should be made on behalf of the user should + // use [UserAuthClient]. + LocalAccessPoint localAccessPoint // tracer is used to create spans tracer oteltrace.Tracer // mfaFuncCache is used to cache the MFA auth method @@ -534,8 +537,8 @@ func (t *CommandHandlerConfig) CheckAndSetDefaults() error { return trace.BadParameter("server: missing server") } - if t.AuthProvider == nil { - return trace.BadParameter("AuthProvider must be provided") + if t.UserAuthClient == nil { + return trace.BadParameter("UserAuthClient must be provided") } if t.SessionCtx == nil { @@ -550,8 +553,8 @@ func (t *CommandHandlerConfig) CheckAndSetDefaults() error { t.TracerProvider = tracing.DefaultProvider() } - if t.LocalAuthProvider == nil { - return trace.BadParameter("LocalAuthProvider must be provided") + if t.LocalAccessPoint == nil { + return trace.BadParameter("localAccessPoint must be provided") } if t.mfaFuncCache == nil { diff --git a/lib/web/terminal.go b/lib/web/terminal.go index e3453ea8145..8aa1563af78 100644 --- a/lib/web/terminal.go +++ b/lib/web/terminal.go @@ -96,9 +96,9 @@ type TerminalRequest struct { ParticipantMode types.SessionParticipantMode `json:"mode"` } -// AuthProvider is a subset of the full Auth API. -type AuthProvider interface { - GetNodes(ctx context.Context, namespace string) ([]types.Server, error) +// UserAuthClient is a subset of the Auth API that performs +// operations on behalf of the user so that the correct RBAC is applied. +type UserAuthClient interface { GetSessionEvents(namespace string, sid session.ID, after int) ([]events.EventFields, error) GetSessionTracker(ctx context.Context, sessionID string) (types.SessionTracker, error) IsMFARequired(ctx context.Context, req *authproto.IsMFARequiredRequest) (*authproto.IsMFARequiredResponse, error) @@ -125,8 +125,8 @@ func NewTerminal(ctx context.Context, cfg TerminalHandlerConfig) (*TerminalHandl "session_id": cfg.SessionData.ID.String(), }), ctx: cfg.SessionCtx, - authProvider: cfg.AuthProvider, - localAuthProvider: cfg.LocalAuthProvider, + userAuthClient: cfg.UserAuthClient, + localAccessPoint: cfg.LocalAccessPoint, sessionData: cfg.SessionData, keepAliveInterval: cfg.KeepAliveInterval, proxyHostPort: cfg.ProxyHostPort, @@ -151,11 +151,14 @@ type TerminalHandlerConfig struct { Term session.TerminalParams // SessionCtx is the context for the users web session. SessionCtx *SessionContext - // AuthProvider is used to fetch nodes and sessions from the backend. - AuthProvider AuthProvider - // LocalAuthProvider is used to fetch user information from the - // local cluster when connecting to agentless nodes. - LocalAuthProvider agentless.AuthProvider + // UserAuthClient is used to fetch nodes and sessions from the backend. + UserAuthClient UserAuthClient + // LocalAccessPoint is the subset of the Proxy cache required to + // look up information from the local cluster. This should not + // be used for anything that requires RBAC on behalf of the user. + // Requests that should be made on behalf of the user should + // use [UserAuthClient]. + LocalAccessPoint localAccessPoint // DisplayLogin is the login name to display in the UI. DisplayLogin string // SessionData is the data to send to the client on the initial session creation. @@ -209,12 +212,12 @@ func (t *TerminalHandlerConfig) CheckAndSetDefaults() error { return trace.BadParameter("term: bad dimensions(%dx%d)", t.Term.W, t.Term.H) } - if t.AuthProvider == nil { - return trace.BadParameter("AuthProvider must be provided") + if t.UserAuthClient == nil { + return trace.BadParameter("UserAuthClient must be provided") } - if t.LocalAuthProvider == nil { - return trace.BadParameter("LocalAuthProvider must be provided") + if t.LocalAccessPoint == nil { + return trace.BadParameter("localAccessPoint must be provided") } if t.SessionCtx == nil { @@ -244,8 +247,8 @@ type sshBaseHandler struct { log *logrus.Entry // ctx is a web session context for the currently logged-in user. ctx *SessionContext - // authProvider is used to fetch nodes and sessions from the backend. - authProvider AuthProvider + // userAuthClient is used to fetch nodes and sessions from the backend via the users' identity. + userAuthClient UserAuthClient // proxyHostPort is the address of the server to connect to. proxyHostPort string // proxyPublicAddr is the public web proxy address. @@ -260,13 +263,24 @@ type sshBaseHandler struct { router *proxy.Router // tracer creates spans tracer oteltrace.Tracer - // localAuthProvider is used to fetch user information from the - // local cluster when connecting to agentless nodes. - localAuthProvider agentless.AuthProvider + // localAccessPoint is the subset of the Proxy cache required to + // look up information from the local cluster. This should not + // be used for anything that requires RBAC on behalf of the user. + // Requests that should be made on behalf of the user should + // use [UserAuthClient]. + localAccessPoint localAccessPoint // interactiveCommand is a command to execute. interactiveCommand []string } +// localAccessPoint is a subset of the cache used to look up +// various cluster details. +type localAccessPoint interface { + GetUser(ctx context.Context, username string, withSecrets bool) (types.User, error) + GetRole(ctx context.Context, name string) (types.Role, error) + GetNode(ctx context.Context, namespace, name string) (types.Server, error) +} + // TerminalHandler connects together an SSH session with a web-based // terminal via a web socket. type TerminalHandler struct { @@ -333,43 +347,60 @@ func (t *TerminalHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } - var sessionMetadataResponse []byte + t.handler(ws, r) +} + +func (t *TerminalHandler) writeSessionData(ctx context.Context) error { + envelope := &Envelope{ + Version: defaults.WebsocketVersion, + Type: defaults.WebsocketSessionMetadata, + } + + sessionDataTemp := t.sessionData // If the displayLogin is set then use it in the session metadata instead of the // login name used in the SSH connection. This is specifically for the use case // when joining a session to avoid displaying "-teleport-internal-join" as the username. if t.displayLogin != "" { - sessionDataTemp := t.sessionData sessionDataTemp.Login = t.displayLogin - sessionMetadataResponse, err = json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp}) + sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp}) + if err != nil { + t.sendError("unable to marshal session response", err, t.stream.ws) + return trace.Wrap(err) + } + envelope.Payload = string(sessionMetadataResponse) } else { - sessionMetadataResponse, err = json.Marshal(siteSessionGenerateResponse{Session: t.sessionData}) - } + // The Proxy cache is used to retrieve the server and resolve the hostname here instead + // of the user auth client to avoid a round trip to the Auth server. This would normally + // not be ok since this bypasses user RBAC, however, since at this point we have already + // established a connection to the target host via the user identity, the user MUST have + // access to the target host. + server, err := t.localAccessPoint.GetNode(ctx, apidefaults.Namespace, sessionDataTemp.ServerID) + if err != nil { + return trace.Wrap(err) + } + sessionDataTemp.ServerHostname = server.GetHostname() - if err != nil { - t.sendError("unable to marshal session response", err, ws) - return - } - - envelope := &Envelope{ - Version: defaults.WebsocketVersion, - Type: defaults.WebsocketSessionMetadata, - Payload: string(sessionMetadataResponse), + sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp}) + if err != nil { + t.sendError("unable to marshal session response", err, t.stream.ws) + return trace.Wrap(err) + } + envelope.Payload = string(sessionMetadataResponse) } envelopeBytes, err := proto.Marshal(envelope) if err != nil { - t.sendError("unable to marshal session data event for web client", err, ws) - return + t.sendError("unable to marshal session data event for web client", err, t.stream.ws) + return trace.Wrap(err) } - err = ws.WriteMessage(websocket.BinaryMessage, envelopeBytes) - if err != nil { - t.sendError("unable to write message to socket", err, ws) - return + if err := t.stream.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes); err != nil { + t.sendError("unable to write message to socket", err, t.stream.ws) + return trace.Wrap(err) } - t.handler(ws, r) + return nil } // Close the websocket stream. @@ -405,7 +436,7 @@ func (t *TerminalHandler) handler(ws *websocket.Conn, r *http.Request) { tctx := oteltrace.ContextWithRemoteSpanContext(context.Background(), oteltrace.SpanContextFromContext(r.Context())) ctx, cancel := context.WithCancel(tctx) defer cancel() - t.stream = NewTerminalStream(ctx, ws, t.log) + t.stream = NewTerminalStream(ctx, TerminalStreamConfig{WS: ws, Logger: t.log}) // Create a Teleport client, if not able to, show the reason to the user in // the terminal. @@ -515,7 +546,7 @@ func (t *TerminalHandler) makeClient(ctx context.Context, stream *TerminalStream // issueSessionMFACerts performs the mfa ceremony to retrieve new certs that can be // used to access nodes which require per-session mfa. The ceremony is performed directly -// to make use of the authProvider already established for the session instead of leveraging +// to make use of the userAuthClient already established for the session instead of leveraging // the TeleportClient which would require dialing the auth server a second time. func (t *sshBaseHandler) issueSessionMFACerts(ctx context.Context, tc *client.TeleportClient, wsStream *WSStream) ([]ssh.AuthMethod, error) { ctx, span := t.tracer.Start(ctx, "terminal/issueSessionMFACerts") @@ -559,7 +590,7 @@ func (t *sshBaseHandler) issueSessionMFACerts(ctx context.Context, tc *client.Te } key, _, err = client.PerformMFACeremony(ctx, client.PerformMFACeremonyParams{ - CurrentAuthClient: t.authProvider, + CurrentAuthClient: t.userAuthClient, RootAuthClient: t.ctx.cfg.RootClient, MFAPrompt: mfa.PromptFunc(func(ctx context.Context, chal *authproto.MFAAuthenticateChallenge) (*authproto.MFAAuthenticateResponse, error) { span.AddEvent("prompting user with mfa challenge") @@ -631,7 +662,7 @@ func (t *sshBaseHandler) connectToHost(ctx context.Context, ws WSConn, tc *clien if err != nil { return nil, trace.Wrap(err) } - signer := agentless.SignerFromSSHCertificate(cert, t.localAuthProvider, tc.SiteName, tc.Username) + signer := agentless.SignerFromSSHCertificate(cert, t.localAccessPoint, tc.SiteName, tc.Username) type clientRes struct { clt *client.NodeClient @@ -751,11 +782,15 @@ func (t *TerminalHandler) streamTerminal(ctx context.Context, tc *client.Telepor defer nc.Close() + if err := t.writeSessionData(ctx); err != nil { + t.log.WithError(err).Warn("Unable to stream terminal - failure sending session data") + } + var beforeStart func(io.Writer) if t.participantMode == types.SessionModeratorMode { beforeStart = func(out io.Writer) { nc.OnMFA = func() { - if err := t.presenceChecker(ctx, out, t.authProvider, t.sessionData.ID.String(), promptMFAChallenge(t.stream.WSStream, protobufMFACodec{})); err != nil { + if err := t.presenceChecker(ctx, out, t.userAuthClient, t.sessionData.ID.String(), promptMFAChallenge(t.stream.WSStream, protobufMFACodec{})); err != nil { t.log.WithError(err).Warn("Unable to stream terminal - failure performing presence checks") return } @@ -968,20 +1003,45 @@ func NewWStream(ctx context.Context, ws WSConn, log logrus.FieldLogger, handlers return w } +// TerminalStreamConfig contains dependencies of a TerminalStream. +type TerminalStreamConfig struct { + // The websocket to operate over. Required. + WS WSConn + // A logger to emit log messages. Optional. + Logger logrus.FieldLogger + // A custom set of handlers to process messages received + // over the websocket. Optional. + Handlers map[string]WSHandlerFunc +} + // NewTerminalStream creates a stream that manages reading and writing // data over the provided [websocket.Conn] -func NewTerminalStream(ctx context.Context, ws WSConn, log logrus.FieldLogger) *TerminalStream { +func NewTerminalStream(ctx context.Context, cfg TerminalStreamConfig) *TerminalStream { t := &TerminalStream{ sessionReadyC: make(chan struct{}), } - handlers := map[string]WSHandlerFunc{ - defaults.WebsocketResize: t.handleWindowResize, - defaults.WebsocketFileTransferRequest: t.handleFileTransferRequest, - defaults.WebsocketFileTransferDecision: t.handleFileTransferDecision, + if cfg.Handlers == nil { + cfg.Handlers = map[string]WSHandlerFunc{} } - t.WSStream = NewWStream(ctx, ws, log, handlers) + if _, ok := cfg.Handlers[defaults.WebsocketResize]; !ok { + cfg.Handlers[defaults.WebsocketResize] = t.handleWindowResize + } + + if _, ok := cfg.Handlers[defaults.WebsocketFileTransferRequest]; !ok { + cfg.Handlers[defaults.WebsocketFileTransferRequest] = t.handleFileTransferRequest + } + + if _, ok := cfg.Handlers[defaults.WebsocketFileTransferDecision]; !ok { + cfg.Handlers[defaults.WebsocketFileTransferDecision] = t.handleFileTransferDecision + } + + if cfg.Logger == nil { + cfg.Logger = utils.NewLogger() + } + + t.WSStream = NewWStream(ctx, cfg.WS, cfg.Logger, cfg.Handlers) return t } diff --git a/lib/web/terminal_test.go b/lib/web/terminal_test.go index 3b222a287a3..5fd0dc1c666 100644 --- a/lib/web/terminal_test.go +++ b/lib/web/terminal_test.go @@ -16,23 +16,32 @@ * along with this program. If not, see . */ -package web_test +package web import ( "context" + "crypto/tls" + "encoding/json" + "fmt" "io" "net/http" "net/http/httptest" + "net/url" "strings" "testing" + "time" "github.com/gogo/protobuf/proto" "github.com/gorilla/websocket" + "github.com/gravitational/roundtrip" + "github.com/gravitational/trace" "github.com/stretchr/testify/require" + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/lib/client" "github.com/gravitational/teleport/lib/defaults" + "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/utils" - "github.com/gravitational/teleport/lib/web" ) // TestTerminalReadFromClosedConn verifies that Teleport recovers @@ -49,7 +58,7 @@ func TestTerminalReadFromClosedConn(t *testing.T) { t.Errorf("couldn't upgrade websocket connection: %v", err) } - envelope := web.Envelope{ + envelope := Envelope{ Type: defaults.WebsocketRaw, Payload: "hello", } @@ -66,7 +75,9 @@ func TestTerminalReadFromClosedConn(t *testing.T) { require.NoError(t, err) defer resp.Body.Close() - stream := web.NewTerminalStream(context.Background(), conn, utils.NewLoggerForTests()) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + stream := NewTerminalStream(ctx, TerminalStreamConfig{WS: conn, Logger: utils.NewLoggerForTests()}) // close the stream before we attempt to read from it, // this will produce a net.ErrClosed error on the read @@ -75,3 +86,164 @@ func TestTerminalReadFromClosedConn(t *testing.T) { _, err = io.Copy(io.Discard, stream) require.NoError(t, err) } + +type terminal struct { + ws *websocket.Conn + stream *TerminalStream + + sessionC chan session.Session +} + +type connectConfig struct { + pack *authPack + host string + proxy string + sessionID session.ID + participantMode types.SessionParticipantMode + keepAliveInterval time.Duration + mfaCeremony func(challenge client.MFAAuthenticateChallenge) []byte + handlers map[string]WSHandlerFunc +} + +func connectToHost(ctx context.Context, cfg connectConfig) (*terminal, error) { + req := TerminalRequest{ + Server: cfg.host, + Login: cfg.pack.login, + Term: session.TerminalParams{ + W: 100, + H: 100, + }, + SessionID: cfg.sessionID, + ParticipantMode: cfg.participantMode, + KeepAliveInterval: cfg.keepAliveInterval, + } + + data, err := json.Marshal(req) + if err != nil { + return nil, trace.Wrap(err) + } + + u := url.URL{ + Host: cfg.proxy, + Scheme: client.WSS, + Path: "/v1/webapi/sites/-current-/connect", + } + + q := u.Query() + q.Set("params", string(data)) + q.Set(roundtrip.AccessTokenQueryParam, cfg.pack.session.Token) + u.RawQuery = q.Encode() + + header := http.Header{} + header.Add("Origin", "http://localhost") + for _, cookie := range cfg.pack.cookies { + header.Add("Cookie", cookie.String()) + } + + dialer := websocket.Dialer{ + TLSClientConfig: &tls.Config{ + InsecureSkipVerify: true, + }, + } + + ws, resp, err := dialer.Dial(u.String(), header) + if err != nil { + var sb strings.Builder + sb.WriteString("websocket dial") + if resp != nil { + fmt.Fprintf(&sb, "; status code %v;", resp.StatusCode) + fmt.Fprintf(&sb, "headers: %v; body: ", resp.Header) + io.Copy(&sb, resp.Body) + } + return nil, trace.Wrap(err, sb.String()) + } + if err := resp.Body.Close(); err != nil { + return nil, trace.Wrap(err) + } + + t := &terminal{ws: ws, sessionC: make(chan session.Session, 1)} + + // If MFA is expected, it should be performed prior to creating + // the TerminalStream to avoid messages being handled by multiple + // readers. + if cfg.mfaCeremony != nil { + if err := t.performMFACeremony(cfg.mfaCeremony); err != nil { + return nil, trace.Wrap(err) + } + } + + if cfg.handlers == nil { + cfg.handlers = map[string]WSHandlerFunc{} + } + + if _, ok := cfg.handlers[defaults.WebsocketSessionMetadata]; !ok { + cfg.handlers[defaults.WebsocketSessionMetadata] = func(ctx context.Context, envelope Envelope) { + if envelope.Type != defaults.WebsocketSessionMetadata { + return + } + + var sessResp siteSessionGenerateResponse + if err := json.Unmarshal([]byte(envelope.Payload), &sessResp); err != nil { + return + } + + t.sessionC <- sessResp.Session + } + } + + t.stream = NewTerminalStream(ctx, TerminalStreamConfig{ + WS: ws, + Logger: utils.NewLogger(), + Handlers: cfg.handlers, + }) + + return t, nil +} + +func (t *terminal) GetSession() session.Session { + sess := <-t.sessionC + t.sessionC <- sess + + return sess +} + +func (t *terminal) Close() error { + return t.stream.Close() +} + +func (t *terminal) Write(p []byte) (int, error) { + return t.stream.Write(p) +} + +func (t *terminal) Read(p []byte) (int, error) { + return t.stream.Read(p) +} + +func (t *terminal) performMFACeremony(ceremonyFn func(challenge client.MFAAuthenticateChallenge) []byte) error { + // Wait for websocket authn challenge event. + ty, raw, err := t.ws.ReadMessage() + if err != nil { + return trace.Wrap(err, "reading ws message") + } + + if ty != websocket.BinaryMessage { + return trace.BadParameter("got unexpected websocket message type %d", ty) + } + + var env Envelope + if err := proto.Unmarshal(raw, &env); err != nil { + return trace.Wrap(err, "unmarshalling envelope") + } + + var challenge client.MFAAuthenticateChallenge + if err := json.Unmarshal([]byte(env.Payload), &challenge); err != nil { + return trace.Wrap(err, "unmarshalling challenge") + } + + // Send response over ws. + if err := t.ws.WriteMessage(websocket.BinaryMessage, ceremonyFn(challenge)); err != nil { + return trace.Wrap(err, "sending challenge response") + } + + return nil +}