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 +}